From ad11f6dcdaba6f01355cb54fd39ab65d1f781d0b Mon Sep 17 00:00:00 2001 From: stuppie Date: Fri, 21 Aug 2026 11:10:06 -0600 Subject: put back a dummy create_bp_payout_event method so I dont break all these tests --- test_utils/models/conftest.py | 6 +----- 1 file changed, 1 insertion(+), 5 deletions(-) (limited to 'test_utils/models') diff --git a/test_utils/models/conftest.py b/test_utils/models/conftest.py index 64bdec6..6a8e4cf 100644 --- a/test_utils/models/conftest.py +++ b/test_utils/models/conftest.py @@ -390,8 +390,6 @@ def bp_payout_factory( amount: Optional["USDCent"] = None, ext_ref_id: Optional[str] = None, created: Optional[AwareDatetime] = None, - skip_wallet_balance_check: bool = False, - skip_one_per_day_check: bool = False, ) -> "BrokerageProductPayoutEvent": from generalresearch.currency import USDCent @@ -402,10 +400,8 @@ def bp_payout_factory( thl_ledger_manager=thl_lm, product=product, amount=amount, - ext_ref_id=ext_ref_id, + ext_ref_id=ext_ref_id or uuid4().hex, created=created, - skip_wallet_balance_check=skip_wallet_balance_check, - skip_one_per_day_check=skip_one_per_day_check, ) return _inner -- cgit v1.2.3 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/conftest.py | 2 +- test_utils/models/contest/conftest.py | 3 ++- test_utils/models/ledger/conftest.py | 2 +- test_utils/models/network/conftest.py | 2 +- 4 files changed, 5 insertions(+), 4 deletions(-) (limited to 'test_utils/models') diff --git a/test_utils/models/conftest.py b/test_utils/models/conftest.py index 89f6f32..93e2f44 100644 --- a/test_utils/models/conftest.py +++ b/test_utils/models/conftest.py @@ -8,8 +8,8 @@ from typing import TYPE_CHECKING, Callable from uuid import uuid4 import pytest -from fastapi import Request from pydantic import AwareDatetime, PositiveInt +from pytest import FixtureRequest as Request from generalresearch.models import Source from generalresearch.models.thl.definitions import ( diff --git a/test_utils/models/contest/conftest.py b/test_utils/models/contest/conftest.py index e750076..bfbe9f8 100644 --- a/test_utils/models/contest/conftest.py +++ b/test_utils/models/contest/conftest.py @@ -6,7 +6,8 @@ from typing import 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.thl.contest_manager import ContestManager 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 diff --git a/test_utils/models/network/conftest.py b/test_utils/models/network/conftest.py index abfbc18..bebc691 100644 --- a/test_utils/models/network/conftest.py +++ b/test_utils/models/network/conftest.py @@ -3,7 +3,7 @@ from datetime import datetime, timedelta, timezone from uuid import uuid4 import pytest -from fastapi import Request +from pytest import FixtureRequest as Request from generalresearch.managers.network.label import IPLabelManager from generalresearch.managers.network.tool_run import ToolRunManager -- 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') 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 e2c5de703be45746bacaea4136f24440ff5a291c Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Mon, 24 Aug 2026 12:35:35 -0700 Subject: Ruff std replacements --- generalresearch/__init__.py | 2 +- generalresearch/config.py | 2 +- generalresearch/grliq/managers/event_plotter.py | 4 +- generalresearch/grliq/managers/forensic_data.py | 4 +- generalresearch/grliq/managers/forensic_events.py | 76 ++++++++++------------ generalresearch/grliq/managers/forensic_results.py | 7 +- generalresearch/grliq/models/custom_types.py | 3 +- generalresearch/grliq/models/forensic_data.py | 1 - generalresearch/grliq/models/forensic_summary.py | 3 +- generalresearch/grliq/utils.py | 2 +- generalresearch/grpc.py | 2 +- generalresearch/incite/base.py | 10 +-- generalresearch/incite/collections/__init__.py | 18 ++--- generalresearch/incite/defaults.py | 2 +- .../incite/mergers/foundations/__init__.py | 2 +- .../incite/mergers/foundations/enriched_session.py | 10 +-- .../mergers/foundations/enriched_task_adjust.py | 5 +- .../incite/mergers/foundations/enriched_wall.py | 8 +-- generalresearch/incite/mergers/ym_survey_wall.py | 10 ++- generalresearch/incite/mergers/ym_wall_summary.py | 3 +- generalresearch/incite/schemas/thl_web.py | 2 +- generalresearch/locales/setup_json.py | 1 + generalresearch/locales/timezone.py | 1 - generalresearch/managers/cint/survey.py | 2 +- generalresearch/managers/criteria.py | 4 +- generalresearch/managers/dynata/survey.py | 2 +- generalresearch/managers/events.py | 4 +- generalresearch/managers/gr/authentication.py | 5 +- generalresearch/managers/gr/team.py | 2 +- generalresearch/managers/innovate/survey.py | 2 +- generalresearch/managers/leaderboard/__init__.py | 3 +- generalresearch/managers/leaderboard/manager.py | 2 +- generalresearch/managers/morning/survey.py | 2 +- generalresearch/managers/network/label.py | 9 ++- generalresearch/managers/network/tool_run.py | 22 +++---- generalresearch/managers/precision/survey.py | 10 +-- generalresearch/managers/prodege/survey.py | 2 +- generalresearch/managers/repdata/survey.py | 4 +- generalresearch/managers/sago/survey.py | 2 +- generalresearch/managers/spectrum/survey.py | 2 +- generalresearch/managers/survey.py | 2 - generalresearch/managers/thl/buyer.py | 2 +- generalresearch/managers/thl/cashout_method.py | 2 +- generalresearch/managers/thl/contest_manager.py | 3 +- .../managers/thl/ledger_manager/conditions.py | 4 +- .../managers/thl/ledger_manager/exceptions.py | 5 -- .../managers/thl/ledger_manager/ledger.py | 15 ++--- .../managers/thl/ledger_manager/thl_ledger.py | 22 +++---- generalresearch/managers/thl/product.py | 10 +-- generalresearch/managers/thl/profiling/question.py | 2 +- generalresearch/managers/thl/profiling/uqa.py | 2 +- generalresearch/managers/thl/profiling/user_upk.py | 2 +- generalresearch/managers/thl/session.py | 11 +--- generalresearch/managers/thl/survey.py | 25 ++++--- generalresearch/managers/thl/task_adjustment.py | 7 +- generalresearch/managers/thl/user_compensate.py | 2 +- .../thl/user_manager/mysql_user_manager.py | 2 +- .../managers/thl/user_manager/user_manager.py | 1 - generalresearch/managers/thl/userhealth.py | 2 +- generalresearch/managers/thl/wall.py | 2 +- generalresearch/managers/thl/wallet/tango.py | 2 +- generalresearch/models/admin/request.py | 2 +- generalresearch/models/custom_types.py | 3 +- generalresearch/models/dynata/survey.py | 4 +- generalresearch/models/gr/authentication.py | 5 +- generalresearch/models/gr/business.py | 6 +- generalresearch/models/gr/team.py | 7 +- generalresearch/models/legacy/questions.py | 3 - generalresearch/models/network/mtr/execute.py | 2 +- generalresearch/models/network/nmap/parser.py | 2 +- generalresearch/models/network/rdns/execute.py | 2 +- generalresearch/models/spectrum/survey.py | 2 +- generalresearch/models/string_utils.py | 1 - generalresearch/models/thl/contest/contest.py | 4 +- .../models/thl/contest/contest_entry.py | 2 +- generalresearch/models/thl/contest/examples.py | 7 -- generalresearch/models/thl/contest/io.py | 2 +- generalresearch/models/thl/contest/leaderboard.py | 2 +- generalresearch/models/thl/contest/milestone.py | 3 +- generalresearch/models/thl/contest/raffle.py | 2 +- generalresearch/models/thl/finance.py | 5 +- generalresearch/models/thl/ipinfo.py | 2 +- generalresearch/models/thl/ledger_example.py | 2 +- generalresearch/models/thl/offerwall/cache.py | 2 +- generalresearch/models/thl/payout_format.py | 12 ++-- generalresearch/models/thl/product.py | 2 - .../models/thl/profiling/marketplace.py | 2 +- .../models/thl/profiling/upk_question.py | 2 +- .../models/thl/profiling/upk_question_answer.py | 2 +- .../models/thl/profiling/user_question_answer.py | 2 +- generalresearch/models/thl/session.py | 2 +- generalresearch/models/thl/survey/__init__.py | 3 - generalresearch/models/thl/survey/buyer.py | 2 +- generalresearch/models/thl/survey/model.py | 2 +- generalresearch/models/thl/survey/penalty.py | 2 +- generalresearch/models/thl/task_adjustment.py | 2 +- generalresearch/models/thl/user.py | 2 +- generalresearch/models/thl/user_iphistory.py | 3 +- generalresearch/models/thl/wallet/payout.py | 2 +- generalresearch/pg_helper.py | 12 ++-- generalresearch/schemas/survey_stats.py | 2 +- generalresearch/sql_helper.py | 19 ++---- generalresearch/thl_django/apps.py | 14 ++-- generalresearch/thl_django/fields.py | 3 +- .../thl_django/migrations/0001_initial.py | 3 +- ..._live_alter_surveycategory_strength_and_more.py | 2 +- ...rveystat_surveystat_live_survey_idx_and_more.py | 2 +- ...ssion_thl_session_status_d578b7_idx_and_more.py | 2 +- ...p_portscanport_iplabel_mtr_portscan_and_more.py | 5 +- generalresearch/thl_django/network/models.py | 5 +- generalresearch/utils/enum.py | 4 +- generalresearch/wall_status_codes/lucid.py | 2 +- generalresearch/wall_status_codes/morning.py | 2 +- generalresearch/wall_status_codes/pollfish.py | 2 +- generalresearch/wall_status_codes/precision.py | 2 +- generalresearch/wall_status_codes/repdata.py | 2 +- test_utils/conftest.py | 3 +- test_utils/grliq/conftest.py | 2 +- test_utils/incite/conftest.py | 2 +- test_utils/managers/conftest.py | 2 - test_utils/models/conftest.py | 9 +-- test_utils/models/contest/conftest.py | 2 +- test_utils/models/network/conftest.py | 2 +- test_utils/models/thl/conftest.py | 2 +- test_utils/spectrum/conftest.py | 2 +- .../incite/collections/test_df_collection_base.py | 3 +- .../collections/test_df_collection_item_base.py | 2 +- .../collections/test_df_collection_item_thl_web.py | 11 +--- .../test_df_collection_thl_marketplaces.py | 7 +- .../collections/test_df_collection_thl_web.py | 2 - .../mergers/foundations/test_enriched_session.py | 7 +- .../foundations/test_enriched_task_adjust.py | 7 -- .../mergers/foundations/test_enriched_wall.py | 23 +------ .../mergers/foundations/test_user_id_product.py | 15 +---- tests/incite/mergers/test_merge_collection.py | 3 +- tests/incite/mergers/test_merge_collection_item.py | 9 +-- tests/incite/mergers/test_pop_ledger.py | 8 +-- tests/incite/mergers/test_ym_survey_merge.py | 17 +---- tests/incite/test_collection_base.py | 3 +- tests/incite/test_collection_base_item.py | 2 +- tests/incite/test_grl_flow.py | 11 ++-- tests/incite/test_interval_idx.py | 3 +- tests/managers/gr/test_authentication.py | 3 - tests/managers/gr/test_business.py | 4 +- tests/managers/gr/test_team.py | 2 - tests/managers/leaderboard.py | 2 +- tests/managers/network/test_label.py | 2 +- tests/managers/test_events.py | 15 ++--- tests/managers/test_userpid.py | 4 +- .../managers/thl/test_contest/test_leaderboard.py | 8 +-- tests/managers/thl/test_contest/test_milestone.py | 14 +--- tests/managers/thl/test_contest/test_raffle.py | 14 +--- tests/managers/thl/test_harmonized_uqa.py | 2 +- tests/managers/thl/test_ipinfo.py | 2 +- tests/managers/thl/test_ledger/test_lm_accounts.py | 6 -- tests/managers/thl/test_ledger/test_lm_tx_locks.py | 23 ++----- .../thl/test_ledger/test_thl_lm_accounts.py | 32 ++++----- .../thl/test_ledger/test_thl_lm_bp_payout.py | 10 +-- tests/managers/thl/test_ledger/test_thl_lm_tx.py | 64 +++--------------- .../test_ledger/test_thl_lm_tx__user_payouts.py | 3 +- tests/managers/thl/test_ledger/test_thl_pem.py | 7 +- tests/managers/thl/test_ledger/test_user_txs.py | 3 +- tests/managers/thl/test_ledger/test_wallet.py | 2 +- tests/managers/thl/test_product.py | 9 ++- tests/managers/thl/test_product_prod.py | 2 - tests/managers/thl/test_profiling/test_user_upk.py | 2 +- tests/managers/thl/test_session_manager.py | 7 +- tests/managers/thl/test_survey.py | 2 +- tests/managers/thl/test_survey_penalty.py | 1 - tests/managers/thl/test_task_adjustment.py | 2 +- tests/managers/thl/test_task_status.py | 2 +- tests/managers/thl/test_user_manager/test_base.py | 4 +- tests/managers/thl/test_user_manager/test_mysql.py | 1 - .../thl/test_user_manager/test_user_fetch.py | 1 - .../thl/test_user_manager/test_user_metadata.py | 1 - tests/managers/thl/test_user_streak.py | 2 +- tests/managers/thl/test_userhealth.py | 2 +- tests/managers/thl/test_wall_manager.py | 13 ++-- tests/models/admin/test_report_request.py | 2 +- tests/models/custom_types/test_aware_datetime.py | 2 +- tests/models/custom_types/test_dsn.py | 1 - tests/models/dynata/test_eligbility.py | 2 +- tests/models/gr/test_authentication.py | 2 +- tests/models/gr/test_business.py | 3 +- tests/models/innovate/test_question.py | 6 +- .../models/legacy/test_user_question_answer_in.py | 4 +- tests/models/morning/test.py | 2 +- tests/models/network/test_mtr.py | 4 +- tests/models/network/test_nmap_parser.py | 1 + tests/models/prodege/test_survey_participation.py | 2 +- tests/models/spectrum/test_question.py | 2 +- tests/models/spectrum/test_survey.py | 2 +- tests/models/spectrum/test_survey_manager.py | 2 +- tests/models/test_finance.py | 2 +- tests/models/thl/question/test_question_info.py | 2 +- tests/models/thl/test_adjustments.py | 30 ++------- .../thl/test_contest/test_leaderboard_contest.py | 2 +- tests/models/thl/test_ledger.py | 2 +- tests/models/thl/test_payout.py | 5 +- tests/models/thl/test_product.py | 2 +- tests/models/thl/test_upkquestion.py | 12 +--- tests/models/thl/test_user.py | 6 +- tests/models/thl/test_user_iphistory.py | 2 +- tests/models/thl/test_user_streak.py | 2 +- tests/models/thl/test_wall.py | 2 +- tests/models/thl/test_wall_session.py | 2 +- tests/sql_helper.py | 4 +- 207 files changed, 398 insertions(+), 725 deletions(-) (limited to 'test_utils/models') diff --git a/generalresearch/__init__.py b/generalresearch/__init__.py index 604b7e2..3b2ec3d 100644 --- a/generalresearch/__init__.py +++ b/generalresearch/__init__.py @@ -129,7 +129,7 @@ def synchronized(wrapped): if lock is None: lock = threading.RLock() - setattr(context, "_synchronized_lock", lock) + context._synchronized_lock = lock return lock diff --git a/generalresearch/config.py b/generalresearch/config.py index 76e3995..c6f41e8 100644 --- a/generalresearch/config.py +++ b/generalresearch/config.py @@ -1,7 +1,7 @@ from __future__ import annotations import os -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from pathlib import Path from pydantic import DirectoryPath, Field, MariaDBDsn, PostgresDsn, RedisDsn diff --git a/generalresearch/grliq/managers/event_plotter.py b/generalresearch/grliq/managers/event_plotter.py index 0bba7c5..94b70ef 100644 --- a/generalresearch/grliq/managers/event_plotter.py +++ b/generalresearch/grliq/managers/event_plotter.py @@ -13,7 +13,7 @@ def make_events_svg( mouse_events: list[MouseEvent], keyboard_events: list[KeyboardEvent] ) -> str: if len(mouse_events) + len(keyboard_events) == 0: - return f'\n' + "\n" + return '\n' + "\n" t = np.array([pm.timeStamp for pm in mouse_events]) t_diff = t.max() - t.min() @@ -88,7 +88,7 @@ def make_events_svg( svg_elements.append(svg_multiline_text(text, cx + 5, cy - 5, font_size)) svg = ( - f'' + '' + "\n".join(svg_elements) + "\n" ) diff --git a/generalresearch/grliq/managers/forensic_data.py b/generalresearch/grliq/managers/forensic_data.py index 7567552..093f7ae 100644 --- a/generalresearch/grliq/managers/forensic_data.py +++ b/generalresearch/grliq/managers/forensic_data.py @@ -1,8 +1,8 @@ from __future__ import annotations +from collections.abc import Collection from datetime import datetime from typing import Any -from collections.abc import Collection from psycopg import sql from pydantic import NonNegativeInt, PositiveInt @@ -611,7 +611,7 @@ class GrlIqDataManager: if res and res["c"] >= 0: return int(res["c"]) - except (Exception,) as e: + except Exception: pass query = f""" diff --git a/generalresearch/grliq/managers/forensic_events.py b/generalresearch/grliq/managers/forensic_events.py index c847d4d..93da481 100644 --- a/generalresearch/grliq/managers/forensic_events.py +++ b/generalresearch/grliq/managers/forensic_events.py @@ -36,36 +36,35 @@ class GrlIqEventManager: "uuid": uuid4().hex, } - with self.postgres_config.make_connection() as conn: - with conn.cursor() as c: - c.execute("SELECT pg_advisory_xact_lock(hashtext(%s))", (session_uuid,)) - # Try to update first - update_query = sql.SQL(""" + with self.postgres_config.make_connection() as conn, conn.cursor() as c: + c.execute("SELECT pg_advisory_xact_lock(hashtext(%s))", (session_uuid,)) + # Try to update first + update_query = sql.SQL(""" UPDATE grliq_forensicevents SET timing_data = %(timing_data)s WHERE session_uuid = %(session_uuid)s AND timing_data IS NULL RETURNING id """) - c.execute(update_query, data) - result = c.fetchone() + c.execute(update_query, data) + result = c.fetchone() - if result: - pk = result["id"] - conn.commit() - return pk + if result: + pk = result["id"] + conn.commit() + return pk - # No matching row to update. Do an insert - insert_query = sql.SQL(""" + # No matching row to update. Do an insert + insert_query = sql.SQL(""" INSERT INTO grliq_forensicevents (uuid, session_uuid, timing_data) VALUES (%(uuid)s, %(session_uuid)s, %(timing_data)s) RETURNING id """) - c.execute(insert_query, data) - pk = c.fetchone()["id"] - conn.commit() + c.execute(insert_query, data) + pk = c.fetchone()["id"] + conn.commit() return int(pk) @@ -88,11 +87,10 @@ class GrlIqEventManager: "event_end": event_end, } - with self.postgres_config.make_connection() as conn: - with conn.cursor() as c: - c.execute("SELECT pg_advisory_xact_lock(hashtext(%s))", (session_uuid,)) - # Try to update first - update_query = sql.SQL(""" + with self.postgres_config.make_connection() as conn, conn.cursor() as c: + c.execute("SELECT pg_advisory_xact_lock(hashtext(%s))", (session_uuid,)) + # Try to update first + update_query = sql.SQL(""" UPDATE grliq_forensicevents SET events = %(events)s, mouse_events = %(mouse_events)s, @@ -102,16 +100,16 @@ class GrlIqEventManager: AND events IS NULL RETURNING id """) - c.execute(update_query, data) - result = c.fetchone() + c.execute(update_query, data) + result = c.fetchone() - if result: - pk = result["id"] - conn.commit() - return pk + if result: + pk = result["id"] + conn.commit() + return pk - # No matching row to update. Do an insert - insert_query = sql.SQL(""" + # No matching row to update. Do an insert + insert_query = sql.SQL(""" INSERT INTO grliq_forensicevents (uuid, session_uuid, events, mouse_events, event_start, event_end) @@ -120,9 +118,9 @@ class GrlIqEventManager: %(event_start)s, %(event_end)s) RETURNING id """) - c.execute(insert_query, data) - pk = c.fetchone()["id"] - conn.commit() + c.execute(insert_query, data) + pk = c.fetchone()["id"] + conn.commit() return int(pk) @@ -167,10 +165,9 @@ class GrlIqEventManager: {filter_str} ORDER BY {order_by} LIMIT {limit} """ - with self.postgres_config.make_connection() as conn: - with conn.cursor() as c: - c.execute(query=query, params=params) - res = c.fetchall() + with self.postgres_config.make_connection() as conn, conn.cursor() as c: + c.execute(query=query, params=params) + res = c.fetchall() for x in res: if x.get("mouse_events"): @@ -206,10 +203,9 @@ class GrlIqEventManager: AND timing_data IS NOT NULL ORDER BY session_uuid, fe.id DESC; """) - with self.postgres_config.make_connection() as conn: - with conn.cursor() as c: - c.execute(query, params) - res = c.fetchall() + with self.postgres_config.make_connection() as conn, conn.cursor() as c: + c.execute(query, params) + res = c.fetchall() for x in res: x["timing_data"] = TimingData.model_validate(x["timing_data"]) diff --git a/generalresearch/grliq/managers/forensic_results.py b/generalresearch/grliq/managers/forensic_results.py index 706a7db..158e582 100644 --- a/generalresearch/grliq/managers/forensic_results.py +++ b/generalresearch/grliq/managers/forensic_results.py @@ -90,10 +90,9 @@ class GrlIqCategoryResultsReader: ORDER BY created_at DESC LIMIT {limit} """ - with self.postgres_config.make_connection() as conn: - with conn.cursor() as c: - c.execute(query, params) - res = c.fetchall() + with self.postgres_config.make_connection() as conn, conn.cursor() as c: + c.execute(query, params) + res = c.fetchall() for x in res: x["client_ip"] = str(x["client_ip"]) diff --git a/generalresearch/grliq/models/custom_types.py b/generalresearch/grliq/models/custom_types.py index 5c7f155..903e230 100644 --- a/generalresearch/grliq/models/custom_types.py +++ b/generalresearch/grliq/models/custom_types.py @@ -1,6 +1,7 @@ -import annotated_types from typing import Annotated +import annotated_types + GrlIqScore = Annotated[int, annotated_types.Ge(0), annotated_types.Le(100)] GrlIqAvgScore = Annotated[float, annotated_types.Ge(0), annotated_types.Le(100)] GrlIqRate = Annotated[float, annotated_types.Ge(0), annotated_types.Le(1)] diff --git a/generalresearch/grliq/models/forensic_data.py b/generalresearch/grliq/models/forensic_data.py index 8d65696..666cb81 100644 --- a/generalresearch/grliq/models/forensic_data.py +++ b/generalresearch/grliq/models/forensic_data.py @@ -779,7 +779,6 @@ class GrlIqData(BaseModel): minutes=90 ), "expired session" - return None def model_dump_sql(self, **kwargs) -> dict[str, Any]: d = dict() diff --git a/generalresearch/grliq/models/forensic_summary.py b/generalresearch/grliq/models/forensic_summary.py index 6f80768..f5b0f25 100644 --- a/generalresearch/grliq/models/forensic_summary.py +++ b/generalresearch/grliq/models/forensic_summary.py @@ -3,7 +3,6 @@ from __future__ import annotations import random from typing import ( Literal, - Optional, Union, get_args, get_origin, @@ -125,7 +124,7 @@ def generate_GrlIqCheckerResultsSummary(): if base_type == GrlIqCheckerResult: if is_opt: fields[f"{field_name}_avg"] = ( - Optional[GrlIqAvgScore], + GrlIqAvgScore | None, Field(default=None, examples=[random.randint(0, 100)]), ) fields[f"{field_name}_pct_none"] = ( diff --git a/generalresearch/grliq/utils.py b/generalresearch/grliq/utils.py index ca8c6a1..bceaa30 100644 --- a/generalresearch/grliq/utils.py +++ b/generalresearch/grliq/utils.py @@ -1,7 +1,7 @@ from __future__ import annotations import os -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime from pathlib import Path from uuid import UUID diff --git a/generalresearch/grpc.py b/generalresearch/grpc.py index 178521e..f1b5611 100644 --- a/generalresearch/grpc.py +++ b/generalresearch/grpc.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from google.protobuf.duration_pb2 import Duration from google.protobuf.timestamp_pb2 import Timestamp diff --git a/generalresearch/incite/base.py b/generalresearch/incite/base.py index 6888ca2..44647fc 100644 --- a/generalresearch/incite/base.py +++ b/generalresearch/incite/base.py @@ -7,8 +7,9 @@ import re import shutil import subprocess import warnings +from collections.abc import Callable, Sequence from concurrent.futures import Future -from datetime import datetime, timedelta, timezone, UTC +from datetime import UTC, datetime, timedelta from os import R_OK, access, listdir from os.path import isdir from os.path import join as pjoin @@ -17,8 +18,8 @@ from sys import platform from typing import ( TYPE_CHECKING, Any, + Self, ) -from collections.abc import Callable, Sequence from uuid import uuid4 import dask @@ -42,7 +43,6 @@ from pydantic import ( ) from pydantic.json_schema import SkipJsonSchema from sentry_sdk import capture_exception -from typing import Self from generalresearch.config import is_debug from generalresearch.incite.schemas import ( @@ -481,7 +481,7 @@ class CollectionBase(BaseModel): item.cleanup_partials() def clear_tmp_archives(self) -> None: - regex = re.compile(r"\.parquet\.[0-9a-f]{32}", re.I) + regex = re.compile(r"\.parquet\.[0-9a-f]{32}", re.IGNORECASE) for fn in os.listdir(self.archive_path): if regex.search(fn): @@ -554,7 +554,7 @@ class CollectionBase(BaseModel): try: pq.ParquetDataset(highest_version).read().to_pandas() - except (Exception,): + except Exception: # If the most recent version isn't valid, we don't want to # create a symlink to it. # TODO: We could try to be smart and iterate down the most recent diff --git a/generalresearch/incite/collections/__init__.py b/generalresearch/incite/collections/__init__.py index 38749b3..42c3d31 100644 --- a/generalresearch/incite/collections/__init__.py +++ b/generalresearch/incite/collections/__init__.py @@ -214,13 +214,13 @@ class DFCollectionItem(CollectionItemBase): """, params=[start, finish], ) - except (Exception,) as e: + except Exception as e: capture_exception(error=e) LOG.error(f"_from_mysql Exception: {e}") return None if not res: - LOG.warning(f"_from_mysql query returned nothing") + LOG.warning("_from_mysql query returned nothing") # Return an empty df.DataFrame with the correct columns return empty_dataframe_from_schema(coll._schema) @@ -228,7 +228,7 @@ class DFCollectionItem(CollectionItemBase): df = self.validate_df(df=df) if df is None: - LOG.warning(f"_from_mysql query results failed validation") + LOG.warning("_from_mysql query results failed validation") # Schema validation can fail... return None @@ -265,13 +265,13 @@ class DFCollectionItem(CollectionItemBase): """, params=[start, finish], ) - except (Exception,) as e: + except Exception as e: capture_exception(error=e) LOG.error(f"_from_postgres Exception: {e}") return None if not res: - LOG.warning(f"_from_postgres query returned nothing") + LOG.warning("_from_postgres query returned nothing") # Return an empty df.DataFrame with the correct columns return empty_dataframe_from_schema(coll._schema) @@ -279,7 +279,7 @@ class DFCollectionItem(CollectionItemBase): df = self.validate_df(df=df) if df is None: - LOG.warning(f"_from_postgres query results failed validation") + LOG.warning("_from_postgres query results failed validation") # Schema validation can fail... return None @@ -351,7 +351,7 @@ class DFCollectionItem(CollectionItemBase): c: Cursor = conn.cursor() for chunk in chunked(tx_ids, n=5_000): c.execute( - query=f""" + query=""" SELECT ltm.transaction_id AS tx_id, ltm.id AS tx_metadata_id, ltm.key, ltm.value @@ -466,7 +466,7 @@ class DFCollectionItem(CollectionItemBase): compression="brotli", ) - except (Exception,) as e: + except Exception as e: LOG.exception(e) self.delete_archive(tmp_path) return False @@ -553,7 +553,7 @@ class DFCollectionItem(CollectionItemBase): write_metadata_file=True, compression="brotli", ) - except (Exception,) as e: + except Exception as e: LOG.exception(e) self.delete_archive(next_numbered_path) return False diff --git a/generalresearch/incite/defaults.py b/generalresearch/incite/defaults.py index 5a95607..d4025fc 100644 --- a/generalresearch/incite/defaults.py +++ b/generalresearch/incite/defaults.py @@ -1,4 +1,4 @@ -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime from generalresearch.incite.base import GRLDatasets from generalresearch.incite.collections import DFCollectionType diff --git a/generalresearch/incite/mergers/foundations/__init__.py b/generalresearch/incite/mergers/foundations/__init__.py index f7a45a8..d3100fa 100644 --- a/generalresearch/incite/mergers/foundations/__init__.py +++ b/generalresearch/incite/mergers/foundations/__init__.py @@ -117,7 +117,7 @@ def annotate_product_and_team_id( try: with conn.cursor() as c: c.execute( - query=f""" + query=""" SELECT u.id AS user_id, u.product_id, bp.team_id FROM thl_user u diff --git a/generalresearch/incite/mergers/foundations/enriched_session.py b/generalresearch/incite/mergers/foundations/enriched_session.py index a368e6c..8a707ca 100644 --- a/generalresearch/incite/mergers/foundations/enriched_session.py +++ b/generalresearch/incite/mergers/foundations/enriched_session.py @@ -60,17 +60,17 @@ class EnrichedSessionMergeItem(MergeCollectionItem): return # --- Session --- - LOG.warning(f"EnrichedSessionMergeItem: get session_collection") + LOG.warning("EnrichedSessionMergeItem: get session_collection") session_items = [w for w in session_coll.items if w.interval.overlaps(ir)] if len(session_items) == 0: - LOG.warning(f"EnrichedSessionMergeItem: no session items. set_empty.") + LOG.warning("EnrichedSessionMergeItem: no session items. set_empty.") if self.should_archive(): self.set_empty() return if not ( session_items[-1].has_partial_archive() or session_items[-1].has_archive() ): - LOG.warning(f"EnrichedSessionMergeItem: session isn't updated!") + LOG.warning("EnrichedSessionMergeItem: session isn't updated!") return sddf = session_coll.ddf( @@ -81,7 +81,7 @@ class EnrichedSessionMergeItem(MergeCollectionItem): ) # --- Walls --- - LOG.warning(f"EnrichedSessionMergeItem: merge wall_collection") + LOG.warning("EnrichedSessionMergeItem: merge wall_collection") wall_items = [ w for w in wall_coll.items @@ -95,7 +95,7 @@ class EnrichedSessionMergeItem(MergeCollectionItem): ] if len(wall_items) == 0: - LOG.error(f"EnrichedSessionMergeItem: no wall items") + LOG.error("EnrichedSessionMergeItem: no wall items") return wddf = wall_coll.ddf( diff --git a/generalresearch/incite/mergers/foundations/enriched_task_adjust.py b/generalresearch/incite/mergers/foundations/enriched_task_adjust.py index f3ab8d8..234cd5b 100644 --- a/generalresearch/incite/mergers/foundations/enriched_task_adjust.py +++ b/generalresearch/incite/mergers/foundations/enriched_task_adjust.py @@ -55,7 +55,7 @@ class EnrichedTaskAdjustMergeItem(MergeCollectionItem): LOG.warning(f"EnrichedReconMergeItem.build({ir})") # --- Task Adjustments --- - LOG.warning(f"EnrichedReconMergeItem: get session_collection") + LOG.warning("EnrichedReconMergeItem: get session_collection") task_adj_coll_items = [ w for w in task_adj_coll.items if w.interval.overlaps(ir) ] @@ -209,6 +209,5 @@ class EnrichedTaskAdjustMerge(MergeCollection): enriched_wall=enriched_wall, pg_config=pg_config, ) - except (Exception,) as e: + except Exception as e: capture_exception(error=e) - pass diff --git a/generalresearch/incite/mergers/foundations/enriched_wall.py b/generalresearch/incite/mergers/foundations/enriched_wall.py index 5a7dd2b..b2ac7bb 100644 --- a/generalresearch/incite/mergers/foundations/enriched_wall.py +++ b/generalresearch/incite/mergers/foundations/enriched_wall.py @@ -55,10 +55,10 @@ class EnrichedWallMergeItem(MergeCollectionItem): return # --- Wall --- - LOG.warning(f"EnrichedWallMergeItem: get wall_collection") + LOG.warning("EnrichedWallMergeItem: get wall_collection") wall_items = [w for w in wall_coll.items if w.interval.overlaps(ir)] if len(wall_items) == 0: - LOG.warning(f"EnrichedWallMergeItem: no wall items. set_empty.") + LOG.warning("EnrichedWallMergeItem: no wall items. set_empty.") if self.should_archive(): self.set_empty() return @@ -93,7 +93,7 @@ class EnrichedWallMergeItem(MergeCollectionItem): wdf = wdf.reset_index(drop=False) # --- Sessions --- - LOG.warning(f"EnrichedWallMergeItem: merge session_collection") + LOG.warning("EnrichedWallMergeItem: merge session_collection") session_items = [ s for s in session_coll.items @@ -107,7 +107,7 @@ class EnrichedWallMergeItem(MergeCollectionItem): ] if len(session_items) == 0: - LOG.error(f"EnrichedWallMergeItem: no session items. breaking early.") + LOG.error("EnrichedWallMergeItem: no session items. breaking early.") return sdf = session_coll.ddf( diff --git a/generalresearch/incite/mergers/ym_survey_wall.py b/generalresearch/incite/mergers/ym_survey_wall.py index c060aae..9750b57 100644 --- a/generalresearch/incite/mergers/ym_survey_wall.py +++ b/generalresearch/incite/mergers/ym_survey_wall.py @@ -62,7 +62,7 @@ class YMSurveyWallMergeCollectionItem(MergeCollectionItem): ) ddf = ddf[ddf["started"] > start] - LOG.warning(f"YMSurveyWallMerge: merge session_collection") + LOG.warning("YMSurveyWallMerge: merge session_collection") session_items = [ s for s in enriched_session.items @@ -98,17 +98,16 @@ class YMSurveyWallMergeCollectionItem(MergeCollectionItem): df.dropna(subset="product_id", how="any", inplace=True) df.sort_values(by="started", inplace=True) - LOG.debug(f"YMSurveyWallMerge.build() validation") + LOG.debug("YMSurveyWallMerge.build() validation") df = self.validate_df(df=df) if df is not None: ddf = dd.from_pandas(df, npartitions=4) - LOG.info(f"YMSurveyWallMerge.build() saving") + LOG.info("YMSurveyWallMerge.build() saving") self.to_archive_symlink(client=client, ddf=ddf) else: LOG.warning("YMSurveyWallMerge failed validation") - return None class YMSurveyWallMerge(MergeCollection): @@ -144,8 +143,7 @@ class YMSurveyWallMerge(MergeCollection): wall_coll=wall_coll, enriched_session=enriched_session, ) - except (Exception,) as e: + except Exception as e: capture_exception(error=e) - pass item.delete_dangling_partials(keep_latest=2, target_path=item.path) diff --git a/generalresearch/incite/mergers/ym_wall_summary.py b/generalresearch/incite/mergers/ym_wall_summary.py index 618b810..37fc3b9 100644 --- a/generalresearch/incite/mergers/ym_wall_summary.py +++ b/generalresearch/incite/mergers/ym_wall_summary.py @@ -117,9 +117,8 @@ class YMWallSummaryMerge(MergeCollection): # item every time build is run even if it isn't closed # if item.should_archive(): item.fetch(wall_collection, session_collection, user_id_product) - except (Exception,) as e: + except Exception as e: capture_exception(e) - pass @staticmethod def build_groupbys(df: pd.DataFrame) -> pd.DataFrame: diff --git a/generalresearch/incite/schemas/thl_web.py b/generalresearch/incite/schemas/thl_web.py index b831b9a..c1be202 100644 --- a/generalresearch/incite/schemas/thl_web.py +++ b/generalresearch/incite/schemas/thl_web.py @@ -1,4 +1,4 @@ -from datetime import datetime, timedelta, timezone, UTC +from datetime import UTC, datetime, timedelta import pandas as pd from pandera.pandas import Check, Column, DataFrameSchema, Index, MultiIndex diff --git a/generalresearch/locales/setup_json.py b/generalresearch/locales/setup_json.py index 356084a..57caa00 100644 --- a/generalresearch/locales/setup_json.py +++ b/generalresearch/locales/setup_json.py @@ -11,6 +11,7 @@ def country_default_lang(): """ raise ValueError("no need to run this, I already ran it.") import pandas as pd + from generalresearch.locales import Localelator l = Localelator() diff --git a/generalresearch/locales/timezone.py b/generalresearch/locales/timezone.py index fce6e0e..810dba3 100644 --- a/generalresearch/locales/timezone.py +++ b/generalresearch/locales/timezone.py @@ -1,4 +1,3 @@ -from typing import Optional from pytz import country_timezones diff --git a/generalresearch/managers/cint/survey.py b/generalresearch/managers/cint/survey.py index da1ecd9..819ae3d 100644 --- a/generalresearch/managers/cint/survey.py +++ b/generalresearch/managers/cint/survey.py @@ -2,7 +2,7 @@ from __future__ import annotations import logging from collections.abc import Collection -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime import pymysql from pymysql import IntegrityError diff --git a/generalresearch/managers/criteria.py b/generalresearch/managers/criteria.py index b5d9830..c13b8ac 100644 --- a/generalresearch/managers/criteria.py +++ b/generalresearch/managers/criteria.py @@ -2,7 +2,7 @@ from __future__ import annotations from abc import ABC from collections.abc import Collection -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from more_itertools import chunked @@ -30,7 +30,6 @@ class CriteriaManager(SqlManager, ABC): """ Create a single criterion """ - ... def filter(self, hashes: Collection[str]) -> dict[str, MarketplaceCondition]: """ @@ -96,7 +95,6 @@ class CriteriaManager(SqlManager, ABC): ) conn.commit() - return None @property def mysql_fields(self) -> str: diff --git a/generalresearch/managers/dynata/survey.py b/generalresearch/managers/dynata/survey.py index 3a15c4d..7643dc4 100644 --- a/generalresearch/managers/dynata/survey.py +++ b/generalresearch/managers/dynata/survey.py @@ -2,7 +2,7 @@ from __future__ import annotations import logging from collections.abc import Collection -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime import pymysql from pymysql import IntegrityError diff --git a/generalresearch/managers/events.py b/generalresearch/managers/events.py index f3c6a04..f6c429e 100644 --- a/generalresearch/managers/events.py +++ b/generalresearch/managers/events.py @@ -5,7 +5,7 @@ import math import socket import threading import time -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal from typing import TYPE_CHECKING @@ -724,7 +724,6 @@ class EventManager(StatsManager): ) ) self.publish_event(msg, product_id=user.product_id) - return def handle_task_finish(self, wall: Wall, session: Session, user: User): self.mark_user_active(user=user) @@ -818,7 +817,6 @@ class EventSubscriber(RedisManager): p.subscribe(self.get_channel_name()) self.pubsub_client = r self.pubsub = p - return def get_channel_name(self): return f"{self.cache_prefix}:event-channel:{self.product_id}" diff --git a/generalresearch/managers/gr/authentication.py b/generalresearch/managers/gr/authentication.py index 409cb10..80bee4b 100644 --- a/generalresearch/managers/gr/authentication.py +++ b/generalresearch/managers/gr/authentication.py @@ -3,7 +3,7 @@ from __future__ import annotations import binascii import logging import os -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime from typing import TYPE_CHECKING, Any from psycopg import sql @@ -146,7 +146,7 @@ class GRUserManager(PostgresManagerWithRedis): for item in res: for k, v in item.items(): - if isinstance(item[k], datetime): + if isinstance(v, datetime): item[k] = item[k].replace(tzinfo=UTC) return [GRUser.model_validate(item) for item in res] @@ -270,7 +270,6 @@ class GRTokenManager(PostgresManager): ) conn.commit() - return def get_by_user_id(self, user_id: PositiveInt) -> GRToken | None: # django authtoken_token table has (user_id) UNIQUE constraint diff --git a/generalresearch/managers/gr/team.py b/generalresearch/managers/gr/team.py index 393f446..d04370b 100644 --- a/generalresearch/managers/gr/team.py +++ b/generalresearch/managers/gr/team.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime from typing import TYPE_CHECKING from uuid import uuid4 diff --git a/generalresearch/managers/innovate/survey.py b/generalresearch/managers/innovate/survey.py index c65b100..cddfba2 100644 --- a/generalresearch/managers/innovate/survey.py +++ b/generalresearch/managers/innovate/survey.py @@ -2,7 +2,7 @@ from __future__ import annotations import logging from collections.abc import Collection -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime import pymysql from pymysql import IntegrityError diff --git a/generalresearch/managers/leaderboard/__init__.py b/generalresearch/managers/leaderboard/__init__.py index d5138cd..048d2cb 100644 --- a/generalresearch/managers/leaderboard/__init__.py +++ b/generalresearch/managers/leaderboard/__init__.py @@ -1,8 +1,9 @@ from __future__ import annotations +from zoneinfo import ZoneInfo + import pytz from cachetools import LRUCache, cached -from zoneinfo import ZoneInfo @cached(cache=LRUCache(maxsize=1)) diff --git a/generalresearch/managers/leaderboard/manager.py b/generalresearch/managers/leaderboard/manager.py index 71c3a73..07e3e2c 100644 --- a/generalresearch/managers/leaderboard/manager.py +++ b/generalresearch/managers/leaderboard/manager.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import datetime, timedelta, timezone, UTC +from datetime import UTC, datetime, timedelta from decimal import Decimal from functools import cached_property from typing import TYPE_CHECKING diff --git a/generalresearch/managers/morning/survey.py b/generalresearch/managers/morning/survey.py index 5fba70d..0e29010 100644 --- a/generalresearch/managers/morning/survey.py +++ b/generalresearch/managers/morning/survey.py @@ -3,7 +3,7 @@ from __future__ import annotations import json import logging from collections.abc import Collection -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime import pymysql from pymysql import IntegrityError diff --git a/generalresearch/managers/network/label.py b/generalresearch/managers/network/label.py index c5306a6..1f44862 100644 --- a/generalresearch/managers/network/label.py +++ b/generalresearch/managers/network/label.py @@ -1,7 +1,7 @@ from __future__ import annotations from collections.abc import Collection -from datetime import datetime, timedelta, timezone, UTC +from datetime import UTC, datetime, timedelta from psycopg import sql from pydantic import IPvAnyNetwork, TypeAdapter @@ -28,10 +28,9 @@ class IPLabelManager(PostgresManager): %(provider)s, %(metadata)s ) RETURNING id;""") params = ip_label.model_dump_postgres() - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute(query, params) - pk = c.fetchone()["id"] + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute(query, params) + pk = c.fetchone()["id"] return ip_label def make_filter_str( diff --git a/generalresearch/managers/network/tool_run.py b/generalresearch/managers/network/tool_run.py index 026b3d3..73afc13 100644 --- a/generalresearch/managers/network/tool_run.py +++ b/generalresearch/managers/network/tool_run.py @@ -49,7 +49,6 @@ class ToolRunManager(PostgresManager): c.execute(query, params) run_id = c.fetchone()["id"] run.id = run_id - return None def create_tool_run(self, run: NmapRun | RDNSRun | MTRRun): if type(run) is NmapRun: @@ -77,10 +76,9 @@ class ToolRunManager(PostgresManager): """ Insert a PortScan + PortScanPorts from a Pydantic NmapResult. """ - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - self._create_tool_run(run, c) - self.nmap_manager._create(run, c=c) + with self.pg_config.make_connection() as conn, conn.cursor() as c: + self._create_tool_run(run, c) + self.nmap_manager._create(run, c=c) return run def get_nmap_run(self, id: int) -> NmapRun: @@ -98,10 +96,9 @@ class ToolRunManager(PostgresManager): """ Insert a RDnsRun + RDNSResult """ - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - self._create_tool_run(run, c) - self.rdns_manager._create(run, c=c) + with self.pg_config.make_connection() as conn, conn.cursor() as c: + self._create_tool_run(run, c) + self.rdns_manager._create(run, c=c) return run def get_rdns_run(self, id: int) -> RDNSRun: @@ -120,10 +117,9 @@ class ToolRunManager(PostgresManager): return RDNSRun.model_validate(res) def create_mtr_run(self, run: MTRRun) -> MTRRun: - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - self._create_tool_run(run, c) - self.mtr_manager._create(run, c=c) + with self.pg_config.make_connection() as conn, conn.cursor() as c: + self._create_tool_run(run, c) + self.mtr_manager._create(run, c=c) return run def get_mtr_run(self, id: int) -> MTRRun: diff --git a/generalresearch/managers/precision/survey.py b/generalresearch/managers/precision/survey.py index 833cb28..c13dca8 100644 --- a/generalresearch/managers/precision/survey.py +++ b/generalresearch/managers/precision/survey.py @@ -2,7 +2,7 @@ from __future__ import annotations import logging from collections.abc import Collection -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime import pymysql from pymysql import IntegrityError @@ -125,7 +125,7 @@ class PrecisionSurveyManager(SurveyManager): country_data = [(survey.survey_id, c) for c in survey.country_isos] c.executemany( - f""" + """ INSERT INTO `thl-precision`.`precision_survey_country` (survey_id, country_iso, is_active) VALUES (%s, %s, TRUE) @@ -134,7 +134,7 @@ class PrecisionSurveyManager(SurveyManager): ) lang_data = [(survey.survey_id, c) for c in survey.language_isos] c.executemany( - f""" + """ INSERT INTO `thl-precision`.`precision_survey_language` (survey_id, language_iso, is_active) VALUES (%s, %s, TRUE) @@ -188,7 +188,7 @@ class PrecisionSurveyManager(SurveyManager): country_data = [(survey.survey_id, c) for c in survey.country_isos] # Turn ON countries in this survey's list of countries, insert row, if already exists, set active. c.executemany( - query=f""" + query=""" INSERT INTO `thl-precision`.`precision_survey_country` (survey_id, country_iso, is_active) VALUES (%s, %s, TRUE) ON DUPLICATE KEY UPDATE is_active = TRUE; @@ -207,7 +207,7 @@ class PrecisionSurveyManager(SurveyManager): ) language_data = [(survey.survey_id, c) for c in survey.language_isos] c.executemany( - query=f""" + query=""" INSERT INTO `thl-precision`.`precision_survey_language` (survey_id, language_iso, is_active) VALUES (%s, %s, TRUE) ON DUPLICATE KEY UPDATE is_active = TRUE; diff --git a/generalresearch/managers/prodege/survey.py b/generalresearch/managers/prodege/survey.py index 750383f..983572c 100644 --- a/generalresearch/managers/prodege/survey.py +++ b/generalresearch/managers/prodege/survey.py @@ -1,7 +1,7 @@ from __future__ import annotations from collections.abc import Collection -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime import pymysql diff --git a/generalresearch/managers/repdata/survey.py b/generalresearch/managers/repdata/survey.py index 1e1c3c6..fe6f621 100644 --- a/generalresearch/managers/repdata/survey.py +++ b/generalresearch/managers/repdata/survey.py @@ -2,7 +2,7 @@ from __future__ import annotations import json from collections.abc import Collection -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime import pymysql @@ -105,7 +105,7 @@ class RepDataSurveyManager(SurveyManager): surveys = {s.survey_id: s for s in surveys} if surveys: res = self.sql_helper.execute_sql_query( - query=f""" + query=""" SELECT * FROM `thl-repdata`.`repdata_surveystream` WHERE survey_id IN %s diff --git a/generalresearch/managers/sago/survey.py b/generalresearch/managers/sago/survey.py index 2582902..462d2ef 100644 --- a/generalresearch/managers/sago/survey.py +++ b/generalresearch/managers/sago/survey.py @@ -2,7 +2,7 @@ from __future__ import annotations import logging from collections.abc import Collection -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime import pymysql from pymysql import IntegrityError diff --git a/generalresearch/managers/spectrum/survey.py b/generalresearch/managers/spectrum/survey.py index 58f8a1a..9b58d43 100644 --- a/generalresearch/managers/spectrum/survey.py +++ b/generalresearch/managers/spectrum/survey.py @@ -2,7 +2,7 @@ from __future__ import annotations import logging from collections.abc import Collection -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime import pymysql from pymysql import IntegrityError diff --git a/generalresearch/managers/survey.py b/generalresearch/managers/survey.py index 964344f..1f057fb 100644 --- a/generalresearch/managers/survey.py +++ b/generalresearch/managers/survey.py @@ -12,14 +12,12 @@ class SurveyManager(SqlManager, ABC): """ Create a single survey """ - ... def update(self, surveys: list[MarketplaceTask]) -> bool: """ Update a list of surveys. Depending on the implementation, this may operate one by one or as a bulk update. """ - ... def update_field(self, survey: MarketplaceTask, field: str) -> bool: """ diff --git a/generalresearch/managers/thl/buyer.py b/generalresearch/managers/thl/buyer.py index ae40bc8..2cb582f 100644 --- a/generalresearch/managers/thl/buyer.py +++ b/generalresearch/managers/thl/buyer.py @@ -1,7 +1,7 @@ from __future__ import annotations from collections.abc import Collection -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime from generalresearch.managers.base import Permission, PostgresManager from generalresearch.models import Source diff --git a/generalresearch/managers/thl/cashout_method.py b/generalresearch/managers/thl/cashout_method.py index 94617f7..10282d4 100644 --- a/generalresearch/managers/thl/cashout_method.py +++ b/generalresearch/managers/thl/cashout_method.py @@ -2,7 +2,7 @@ from __future__ import annotations from collections.abc import Collection from copy import copy -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime from typing import Any from uuid import UUID, uuid4 diff --git a/generalresearch/managers/thl/contest_manager.py b/generalresearch/managers/thl/contest_manager.py index 286de3d..62146d7 100644 --- a/generalresearch/managers/thl/contest_manager.py +++ b/generalresearch/managers/thl/contest_manager.py @@ -1,7 +1,7 @@ from __future__ import annotations from collections.abc import Collection -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime from typing import Any, Literal, cast from uuid import UUID @@ -827,7 +827,6 @@ class MilestoneContestManager(ContestBaseManager): ) self.end_milestone_contest(contest) - return None def enter_contest_db_work_milestone( self, contest: MilestoneUserView, user: User, incr: PositiveInt diff --git a/generalresearch/managers/thl/ledger_manager/conditions.py b/generalresearch/managers/thl/ledger_manager/conditions.py index f79a0f0..b2fd465 100644 --- a/generalresearch/managers/thl/ledger_manager/conditions.py +++ b/generalresearch/managers/thl/ledger_manager/conditions.py @@ -1,9 +1,9 @@ from __future__ import annotations import logging -from datetime import datetime, timedelta, timezone, UTC -from typing import TYPE_CHECKING from collections.abc import Callable +from datetime import UTC, datetime, timedelta +from typing import TYPE_CHECKING from generalresearch.config import JAMES_BILLINGS_BPID, JAMES_BILLINGS_TX_CUTOFF from generalresearch.currency import USDCent diff --git a/generalresearch/managers/thl/ledger_manager/exceptions.py b/generalresearch/managers/thl/ledger_manager/exceptions.py index b79c153..48c102b 100644 --- a/generalresearch/managers/thl/ledger_manager/exceptions.py +++ b/generalresearch/managers/thl/ledger_manager/exceptions.py @@ -11,7 +11,6 @@ class LedgerTransactionCreateError(Exception): Ledger transaction creation failed """ - pass class LedgerTransactionCreateLockError(LedgerTransactionCreateError): @@ -19,7 +18,6 @@ class LedgerTransactionCreateLockError(LedgerTransactionCreateError): Ledger transaction creation failed because we could not acquire a lock """ - pass class LedgerTransactionReleaseLockError(LedgerTransactionCreateError): @@ -29,7 +27,6 @@ class LedgerTransactionReleaseLockError(LedgerTransactionCreateError): back-populate as in sentry I see this very rarely. """ - pass class LedgerTransactionFlagAlreadyExistsError(LedgerTransactionCreateError): @@ -38,7 +35,6 @@ class LedgerTransactionFlagAlreadyExistsError(LedgerTransactionCreateError): tx was already set """ - pass class LedgerTransactionConditionFailedError(LedgerTransactionCreateError): @@ -46,4 +42,3 @@ class LedgerTransactionConditionFailedError(LedgerTransactionCreateError): We tried to create a transaction but the condition check failed. """ - pass diff --git a/generalresearch/managers/thl/ledger_manager/ledger.py b/generalresearch/managers/thl/ledger_manager/ledger.py index 00fac27..410f6ca 100644 --- a/generalresearch/managers/thl/ledger_manager/ledger.py +++ b/generalresearch/managers/thl/ledger_manager/ledger.py @@ -2,10 +2,9 @@ from __future__ import annotations import logging from collections import defaultdict -from collections.abc import Collection -from datetime import datetime, timedelta, timezone, UTC +from collections.abc import Callable, Collection +from datetime import UTC, datetime, timedelta from typing import Any -from collections.abc import Callable from uuid import UUID import redis @@ -343,7 +342,7 @@ class LedgerTransactionManager(LedgerManagerBasePostgres): assert len(tag) > 6, "Please confirm the tag is valid" res = self.pg_config.execute_sql_query( - query=f""" + query=""" SELECT lt.id FROM ledger_transaction AS lt WHERE tag = %s @@ -361,7 +360,7 @@ class LedgerTransactionManager(LedgerManagerBasePostgres): def get_tx_ids_by_tags(self, tags: list[str]) -> set[PositiveInt]: res = self.pg_config.execute_sql_query( - query=f""" + query=""" SELECT lt.id, lt.tag, lt.created, lt.ext_description FROM ledger_transaction AS lt WHERE tag = ANY(%s) @@ -869,7 +868,7 @@ class LedgerAccountManager(LedgerManagerBasePostgres): # qualified_name has a unique index so there can only be 0 or 1 match. res = self.pg_config.execute_sql_query( - query=f""" + query=""" SELECT uuid, display_name, qualified_name, account_type, normal_balance, reference_type, @@ -928,7 +927,7 @@ class LedgerAccountManager(LedgerManagerBasePostgres): # TODO: Move to RR with long timeout (2min+), it causes problems res = self.pg_config.execute_sql_query( - query=f""" + query=""" SELECT SUM(amount * direction) AS total FROM ledger_entry WHERE account_id = %s @@ -1046,7 +1045,7 @@ class LedgerManager( """This is for testing only, as it'll take forever to run this if the ledger_manager is huge """ - res = self.pg_config.execute_sql_query(f""" + res = self.pg_config.execute_sql_query(""" SELECT SUM(CASE WHEN normal_balance = -1 THEN total ELSE 0 END) AS credit_total, SUM(CASE WHEN normal_balance = 1 THEN total ELSE 0 END) AS debit_total diff --git a/generalresearch/managers/thl/ledger_manager/thl_ledger.py b/generalresearch/managers/thl/ledger_manager/thl_ledger.py index c9203f6..ded6518 100644 --- a/generalresearch/managers/thl/ledger_manager/thl_ledger.py +++ b/generalresearch/managers/thl/ledger_manager/thl_ledger.py @@ -1,11 +1,10 @@ from __future__ import annotations import logging -from collections.abc import Collection -from datetime import datetime, timedelta, timezone, UTC +from collections.abc import Callable, Collection +from datetime import UTC, datetime, timedelta from decimal import Decimal from typing import TYPE_CHECKING -from collections.abc import Callable from uuid import UUID import numpy as np @@ -407,8 +406,7 @@ class ThlLedgerManager(LedgerManager): f"bp_pay {bp_pay} > thl_net {thl_net}. Capping bp_pay to thl_net." ) bp_pay = thl_net - if user_pay > bp_pay: - user_pay = bp_pay + user_pay = min(user_pay, bp_pay) commission_amount = round(thl_net - bp_pay) @@ -538,7 +536,7 @@ class ThlLedgerManager(LedgerManager): ] else: - logger.info(f"create_transaction_task_adjustment. No transactions needed.") + logger.info("create_transaction_task_adjustment. No transactions needed.") return None amt_str = f"${abs(change_amount) / 100:,.2f}" @@ -666,7 +664,7 @@ class ThlLedgerManager(LedgerManager): else: logger.info( - f"create_transaction_bp_adjustment. No transactions needed." + "create_transaction_bp_adjustment. No transactions needed." ) return None else: @@ -738,7 +736,7 @@ class ThlLedgerManager(LedgerManager): else: logger.info( - f"create_transaction_bp_adjustment. No transactions needed." + "create_transaction_bp_adjustment. No transactions needed." ) return None @@ -856,7 +854,7 @@ class ThlLedgerManager(LedgerManager): ), ] - ext_description = f"BP Payout" + ext_description = "BP Payout" t = self.create_tx( entries=entries, metadata=metadata, @@ -983,7 +981,7 @@ class ThlLedgerManager(LedgerManager): raise ValueError("Invalid Direction") if description is None: - description = f"BP Plug" + description = "BP Plug" t = self.create_tx( entries=entries, @@ -1191,7 +1189,7 @@ class ThlLedgerManager(LedgerManager): f"Trying to cancel user payout {payout_event.uuid} with no request tx found." ) - description = f"User Payout Cancelled" + description = "User Payout Cancelled" f = lambda: self.create_tx_user_payout_cancelled_( user=user, payout_event=payout_event, @@ -1900,7 +1898,7 @@ class ThlLedgerManager(LedgerManager): reserve = round(wall["user_payout_int"].sum() - wall["redeemable"].sum()) redeemable_balance = user_wallet_balance - reserve - redeemable_balance = 0 if redeemable_balance < 0 else redeemable_balance + redeemable_balance = max(redeemable_balance, 0) if redeemable_balance > 0: # it is possible the user_wallet_balance is negative, in which case diff --git a/generalresearch/managers/thl/product.py b/generalresearch/managers/thl/product.py index 3b92361..d924e17 100644 --- a/generalresearch/managers/thl/product.py +++ b/generalresearch/managers/thl/product.py @@ -4,7 +4,7 @@ import json import logging import operator from collections.abc import Collection -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime from decimal import Decimal from threading import Lock from typing import TYPE_CHECKING @@ -94,9 +94,9 @@ class ProductManager(PostgresManager): return self.fetch_uuids( product_uuids=[product_uuid], )[0] - except (AssertionError,): + except AssertionError: return None - except (IndexError,): + except IndexError: return None def get_by_uuids_if_exists( @@ -397,11 +397,11 @@ class ProductManager(PostgresManager): # # from pymysql import IntegrityError # except IntegrityError as e: - except (Exception,) as e: + except Exception as e: try: return self.get_by_uuid(product_uuid=instance.id) - except (Exception,) as e2: + except Exception: pass finally: self.cache_clear(instance.id) diff --git a/generalresearch/managers/thl/profiling/question.py b/generalresearch/managers/thl/profiling/question.py index 10a9e32..078894a 100644 --- a/generalresearch/managers/thl/profiling/question.py +++ b/generalresearch/managers/thl/profiling/question.py @@ -85,7 +85,7 @@ class QuestionManager(PostgresManager): def lookup_by_property( self, property_code: str, country_iso: str, language_iso: str ) -> UpkQuestion: - query = f""" + query = """ SELECT data, property_code, explanation_template, explanation_fragment_template FROM marketplace_question WHERE property_code = %(property_code)s diff --git a/generalresearch/managers/thl/profiling/uqa.py b/generalresearch/managers/thl/profiling/uqa.py index 3854333..fa1747b 100644 --- a/generalresearch/managers/thl/profiling/uqa.py +++ b/generalresearch/managers/thl/profiling/uqa.py @@ -2,7 +2,7 @@ from __future__ import annotations import logging from collections.abc import Collection -from datetime import datetime, timedelta, timezone, UTC +from datetime import UTC, datetime, timedelta from generalresearch.managers.base import PostgresManagerWithRedis from generalresearch.models.thl.profiling.user_question_answer import ( diff --git a/generalresearch/managers/thl/profiling/user_upk.py b/generalresearch/managers/thl/profiling/user_upk.py index ee36124..a2cddb3 100644 --- a/generalresearch/managers/thl/profiling/user_upk.py +++ b/generalresearch/managers/thl/profiling/user_upk.py @@ -3,7 +3,7 @@ from __future__ import annotations import json from collections import defaultdict from collections.abc import Collection -from datetime import datetime, timedelta, timezone, UTC +from datetime import UTC, datetime, timedelta from typing import Any from uuid import UUID diff --git a/generalresearch/managers/thl/session.py b/generalresearch/managers/thl/session.py index 771f882..959003c 100644 --- a/generalresearch/managers/thl/session.py +++ b/generalresearch/managers/thl/session.py @@ -1,7 +1,7 @@ from __future__ import annotations from collections.abc import Collection -from datetime import datetime, timedelta, timezone, UTC +from datetime import UTC, datetime, timedelta from decimal import Decimal from typing import Any from uuid import UUID, uuid4 @@ -190,14 +190,7 @@ class SessionManager(PostgresManager): # re-run model_validate after finished = finished if finished else datetime.now(tz=UTC) session.update( - **{ - "status": status, - "status_code_1": status_code_1, - "status_code_2": status_code_2, - "finished": finished, - "payout": payout, - "user_payout": user_payout, - } + status=status, status_code_1=status_code_1, status_code_2=status_code_2, finished=finished, payout=payout, user_payout=user_payout ) d = session.model_dump_mysql() self.pg_config.execute_write( diff --git a/generalresearch/managers/thl/survey.py b/generalresearch/managers/thl/survey.py index 966e96d..a9ec841 100644 --- a/generalresearch/managers/thl/survey.py +++ b/generalresearch/managers/thl/survey.py @@ -2,7 +2,7 @@ from __future__ import annotations from collections import defaultdict from collections.abc import Collection -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from typing import Any import pandas as pd @@ -280,7 +280,6 @@ class SurveyManager(PostgresManager): query=query, params={"survey_pks": survey_pks}, ) - return None def update_surveys_categories(self, surveys: list[Survey] | None = None) -> None: for chunk in chunked(surveys, 500): @@ -329,12 +328,11 @@ class SurveyManager(PostgresManager): ] with self.pg_config.make_connection() as conn: # noinspection PyArgumentList - with conn.transaction(): - with conn.cursor() as c: - c.execute(temp_table_sql) - c.executemany(insert_values_sql, rows) - c.execute(delete_sql) - c.execute(upsert_sql) + with conn.transaction(), conn.cursor() as c: + c.execute(temp_table_sql) + c.executemany(insert_values_sql, rows) + c.execute(delete_sql) + c.execute(upsert_sql) conn.commit() def get_survey_categories(self): @@ -760,12 +758,11 @@ class SurveyStatManager(PostgresManager): print(query) print(params) - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute("SET work_mem = '256MB';") - c.execute("SET statement_timeout = '10s';") - c.execute(query, params=params) - res = c.fetchall() + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute("SET work_mem = '256MB';") + c.execute("SET statement_timeout = '10s';") + c.execute(query, params=params) + res = c.fetchall() return [SurveyStat.model_validate(x) for x in res] diff --git a/generalresearch/managers/thl/task_adjustment.py b/generalresearch/managers/thl/task_adjustment.py index 3ec3d41..60bade7 100644 --- a/generalresearch/managers/thl/task_adjustment.py +++ b/generalresearch/managers/thl/task_adjustment.py @@ -1,7 +1,7 @@ from __future__ import annotations import logging -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from decimal import Decimal from functools import cached_property @@ -148,10 +148,7 @@ class TaskAdjustmentManager(PostgresManager): if ( wall.status == Status.COMPLETE and adjusted_status == WallAdjustedStatus.ADJUSTED_TO_COMPLETE - ): - new_adjusted_status = None - new_adjusted_cpi = None - elif ( + ) or ( wall.status != Status.COMPLETE and adjusted_status == WallAdjustedStatus.ADJUSTED_TO_FAIL ): diff --git a/generalresearch/managers/thl/user_compensate.py b/generalresearch/managers/thl/user_compensate.py index c6c0747..8338424 100644 --- a/generalresearch/managers/thl/user_compensate.py +++ b/generalresearch/managers/thl/user_compensate.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from decimal import Decimal from uuid import uuid4 diff --git a/generalresearch/managers/thl/user_manager/mysql_user_manager.py b/generalresearch/managers/thl/user_manager/mysql_user_manager.py index 7931ba4..d2d0ffc 100644 --- a/generalresearch/managers/thl/user_manager/mysql_user_manager.py +++ b/generalresearch/managers/thl/user_manager/mysql_user_manager.py @@ -2,7 +2,7 @@ from __future__ import annotations import logging from collections.abc import Collection -from datetime import datetime, timezone, UTC +from datetime import UTC, datetime from functools import lru_cache from uuid import uuid4 diff --git a/generalresearch/managers/thl/user_manager/user_manager.py b/generalresearch/managers/thl/user_manager/user_manager.py index 3794020..7486869 100644 --- a/generalresearch/managers/thl/user_manager/user_manager.py +++ b/generalresearch/managers/thl/user_manager/user_manager.py @@ -103,7 +103,6 @@ class UserManager: event_value=event_value, ) - return None def cache_clear(self): # Generally this is used in testing. This clears the .get_user's lru_cache. diff --git a/generalresearch/managers/thl/userhealth.py b/generalresearch/managers/thl/userhealth.py index 0bc60ec..babed04 100644 --- a/generalresearch/managers/thl/userhealth.py +++ b/generalresearch/managers/thl/userhealth.py @@ -2,7 +2,7 @@ from __future__ import annotations import ipaddress from collections.abc import Collection -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from itertools import zip_longest from typing import Any diff --git a/generalresearch/managers/thl/wall.py b/generalresearch/managers/thl/wall.py index 7e413d7..03ca1c6 100644 --- a/generalresearch/managers/thl/wall.py +++ b/generalresearch/managers/thl/wall.py @@ -3,7 +3,7 @@ from __future__ import annotations import logging from collections import defaultdict from collections.abc import Collection -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal from functools import cached_property from uuid import uuid4 diff --git a/generalresearch/managers/thl/wallet/tango.py b/generalresearch/managers/thl/wallet/tango.py index a4ebf22..4abfc70 100644 --- a/generalresearch/managers/thl/wallet/tango.py +++ b/generalresearch/managers/thl/wallet/tango.py @@ -44,7 +44,7 @@ def complete_tango_order( tango_client=tango_client, ) - except Exception as e: + except Exception: # todo: its possible the order went through, but something else was wrong # we should try to retrieve the order by its ref_id and confirm it really # failed... diff --git a/generalresearch/models/admin/request.py b/generalresearch/models/admin/request.py index 2d68de1..5fdc784 100644 --- a/generalresearch/models/admin/request.py +++ b/generalresearch/models/admin/request.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from enum import Enum from typing import Literal diff --git a/generalresearch/models/custom_types.py b/generalresearch/models/custom_types.py index 9346064..c200b34 100644 --- a/generalresearch/models/custom_types.py +++ b/generalresearch/models/custom_types.py @@ -2,8 +2,7 @@ from __future__ import annotations import json import re -import sys as _sys -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from typing import Annotated, Any, Literal from uuid import UUID diff --git a/generalresearch/models/dynata/survey.py b/generalresearch/models/dynata/survey.py index 9adac42..5a9f763 100644 --- a/generalresearch/models/dynata/survey.py +++ b/generalresearch/models/dynata/survey.py @@ -132,9 +132,7 @@ class DynataCondition(MarketplaceCondition): if cell["kind"] == "RANGE": d["values"] = [ - "{0}-{1}".format( - cell["range"]["from"] or "inf", cell["range"]["to"] or "inf" - ) + f"{cell["range"]["from"] or "inf"}-{cell["range"]["to"] or "inf"}" ] d["value_type"] = ConditionValueType.RANGE return cls.model_validate(d) diff --git a/generalresearch/models/gr/authentication.py b/generalresearch/models/gr/authentication.py index 67a8fc2..f9644fe 100644 --- a/generalresearch/models/gr/authentication.py +++ b/generalresearch/models/gr/authentication.py @@ -3,7 +3,7 @@ from __future__ import annotations import binascii import json import os -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from typing import TYPE_CHECKING, Any, Self from pydantic import ( @@ -181,7 +181,7 @@ class GRUser(BaseModel): team_products = pm.fetch_uuids(team_uuids=team_uuids) if team_uuids else [] products = {p.id: p for p in business_products + team_products} - self.products = sorted(products.values(), key=lambda x: getattr(x, "created")) + self.products = sorted(products.values(), key=lambda x: x.created) def prefetch_token(self, pg_config: PostgresConfig): from generalresearch.managers.gr.authentication import ( @@ -283,7 +283,6 @@ class GRUser(BaseModel): ex=ex_secs, ) - return None # --- ORM --- diff --git a/generalresearch/models/gr/business.py b/generalresearch/models/gr/business.py index 534b23f..74b5c29 100644 --- a/generalresearch/models/gr/business.py +++ b/generalresearch/models/gr/business.py @@ -420,7 +420,7 @@ class Business(BaseModel): # that is still valid. Don't attempt to build a balance, leave it # as None rather than all zeros LOG.warning(f"Business({self.uuid=}).prebuild_balance empty dataframe") - return None + return LOG.debug(f"Business.prebuild_balance.groupby() {df.head()}") df = df.groupby("account_id").sum() @@ -442,13 +442,11 @@ class Business(BaseModel): order_by=OrderBy.DESC, ) self.prebuild_payouts_total() - return None def prebuild_payouts_total(self): assert self.payouts is not None self.payouts_total = USDCent(sum([po.amount for po in self.payouts])) self.payouts_total_str = self.payouts_total.to_usd_str() - return None def prebuild_pop_financial( self, @@ -548,7 +546,6 @@ class Business(BaseModel): except Exception as e: raise OSError(f"Parquet verification failed: {e}") - return None def prebuild_enriched_wall_parquet( self, @@ -593,7 +590,6 @@ class Business(BaseModel): except Exception as e: raise OSError(f"Parquet verification failed: {e}") - return None @classmethod def required_fields(cls) -> list[str]: diff --git a/generalresearch/models/gr/team.py b/generalresearch/models/gr/team.py index 900062f..78a9ba9 100644 --- a/generalresearch/models/gr/team.py +++ b/generalresearch/models/gr/team.py @@ -2,7 +2,7 @@ from __future__ import annotations import json import os -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from enum import Enum from pathlib import Path from typing import TYPE_CHECKING, Self @@ -194,7 +194,6 @@ class Team(BaseModel): except Exception as e: raise OSError(f"Parquet verification failed: {e}") - return def prebuild_enriched_wall_parquet( self, @@ -239,7 +238,6 @@ class Team(BaseModel): except Exception as e: raise OSError(f"Parquet verification failed: {e}") - return None @classmethod def required_fields(cls) -> list[str]: @@ -325,7 +323,6 @@ class Team(BaseModel): enriched_wall=enriched_wall, ) - return # --- ORM --- @@ -344,5 +341,5 @@ class Team(BaseModel): d = {val: json.loads(res[idx]) for idx, val in enumerate(keys)} return Team.model_validate(d) - except (Exception,) as e: + except Exception: return None diff --git a/generalresearch/models/legacy/questions.py b/generalresearch/models/legacy/questions.py index 8e19e57..4651ab0 100644 --- a/generalresearch/models/legacy/questions.py +++ b/generalresearch/models/legacy/questions.py @@ -106,7 +106,6 @@ class UserQuestionAnswerIn(BaseModel): if self.question_id == user_agent_qid: val = self.answer[0] # assert val == request.user_agent.to_header(): - pass return self @@ -217,7 +216,6 @@ class UserQuestionAnswers(BaseModel): # --- Prefetch --- def prefetch_user(self, um: UserManager) -> None: - from generalresearch.models.thl.user import User res: User | None = um.get_user_if_exists( product_id=self.product_id, product_user_id=self.product_user_id @@ -230,7 +228,6 @@ class UserQuestionAnswers(BaseModel): def prefetch_wall(self, wm: WallManager) -> None: from generalresearch.models import Source - from generalresearch.models.thl.session import Wall res: Wall | None = wm.get_from_uuid_if_exists(wall_uuid=self.session_id) diff --git a/generalresearch/models/network/mtr/execute.py b/generalresearch/models/network/mtr/execute.py index 953124d..c5b3c5c 100644 --- a/generalresearch/models/network/mtr/execute.py +++ b/generalresearch/models/network/mtr/execute.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from uuid import uuid4 from generalresearch.models.custom_types import UUIDStr diff --git a/generalresearch/models/network/nmap/parser.py b/generalresearch/models/network/nmap/parser.py index 6ad4ab4..ecaf2d1 100644 --- a/generalresearch/models/network/nmap/parser.py +++ b/generalresearch/models/network/nmap/parser.py @@ -1,7 +1,7 @@ from __future__ import annotations import xml.etree.ElementTree as ET -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from typing import Any from generalresearch.models.network.definitions import IPProtocol diff --git a/generalresearch/models/network/rdns/execute.py b/generalresearch/models/network/rdns/execute.py index 1d74df2..d6de84b 100644 --- a/generalresearch/models/network/rdns/execute.py +++ b/generalresearch/models/network/rdns/execute.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from uuid import uuid4 from generalresearch.models.custom_types import UUIDStr diff --git a/generalresearch/models/spectrum/survey.py b/generalresearch/models/spectrum/survey.py index aaf182e..f9a5e27 100644 --- a/generalresearch/models/spectrum/survey.py +++ b/generalresearch/models/spectrum/survey.py @@ -75,7 +75,7 @@ class SpectrumCondition(MarketplaceCondition): rs["from"] = round(rs["from"] / 12) rs["to"] = round(rs["to"] / 12) d["values"] = [ - "{0}-{1}".format(rs["from"] or "inf", rs["to"] or "inf") + f"{rs["from"] or "inf"}-{rs["to"] or "inf"}" for rs in d["range_sets"] ] d["value_type"] = ConditionValueType.RANGE diff --git a/generalresearch/models/string_utils.py b/generalresearch/models/string_utils.py index d76456f..dff2f4d 100644 --- a/generalresearch/models/string_utils.py +++ b/generalresearch/models/string_utils.py @@ -1,5 +1,4 @@ import unicodedata -from typing import Optional def remove_nbsp(s: str | None) -> str | None: diff --git a/generalresearch/models/thl/contest/contest.py b/generalresearch/models/thl/contest/contest.py index 8c45f13..2a8853d 100644 --- a/generalresearch/models/thl/contest/contest.py +++ b/generalresearch/models/thl/contest/contest.py @@ -2,7 +2,7 @@ from __future__ import annotations import json from abc import ABC, abstractmethod -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from typing import Any, Self from uuid import uuid4 @@ -167,7 +167,7 @@ class Contest(ContestBase): ended_at=datetime.now(tz=UTC), end_reason=reason, ) - return None + return def model_dump_mysql(self, **kwargs) -> dict[str, Any]: d = self.model_dump(mode="json", **kwargs) diff --git a/generalresearch/models/thl/contest/contest_entry.py b/generalresearch/models/thl/contest/contest_entry.py index 31ef317..bb3aef4 100644 --- a/generalresearch/models/thl/contest/contest_entry.py +++ b/generalresearch/models/thl/contest/contest_entry.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from uuid import uuid4 from pydantic import ( diff --git a/generalresearch/models/thl/contest/examples.py b/generalresearch/models/thl/contest/examples.py index 1810e63..018f090 100644 --- a/generalresearch/models/thl/contest/examples.py +++ b/generalresearch/models/thl/contest/examples.py @@ -93,7 +93,6 @@ def _example_raffle(schema: dict) -> None: product_id=EXAMPLE_PRODUCT_ID, ).model_dump(mode="json") - return None def _example_raffle_user_view(schema: dict[str, Any]) -> None: @@ -144,7 +143,6 @@ def _example_raffle_user_view(schema: dict[str, Any]) -> None: product_user_id="test-user", ).model_dump(mode="json") - return None def _example_milestone_create(schema: dict[str, Any]) -> None: @@ -185,7 +183,6 @@ def _example_milestone_create(schema: dict[str, Any]) -> None: terms_and_conditions=HttpUrl("https://www.example.com"), ).model_dump(mode="json") - return None def _example_milestone(schema: dict[str, Any]) -> None: @@ -231,7 +228,6 @@ def _example_milestone(schema: dict[str, Any]) -> None: win_count=12, ).model_dump(mode="json") - return None def _example_milestone_user_view(schema: dict[str, Any]) -> None: @@ -277,7 +273,6 @@ def _example_milestone_user_view(schema: dict[str, Any]) -> None: product_user_id="test-user", ).model_dump(mode="json") - return None def _example_leaderboard_contest_create(schema: dict[str, Any]) -> None: @@ -322,7 +317,6 @@ def _example_leaderboard_contest_create(schema: dict[str, Any]) -> None: leaderboard_key=f"leaderboard:{EXAMPLE_PRODUCT_ID}:us:weekly:2025-05-26:complete_count", ).model_dump(mode="json") - return None def _example_leaderboard_contest(schema: dict[str, Any]) -> None: @@ -368,7 +362,6 @@ def _example_leaderboard_contest(schema: dict[str, Any]) -> None: product_id=EXAMPLE_PRODUCT_ID, ).model_dump(mode="json") - return None def _example_leaderboard_contest_user_view(schema: dict[str, Any]) -> None: diff --git a/generalresearch/models/thl/contest/io.py b/generalresearch/models/thl/contest/io.py index c6af719..d11080f 100644 --- a/generalresearch/models/thl/contest/io.py +++ b/generalresearch/models/thl/contest/io.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from uuid import uuid4 from generalresearch.models.thl.contest.definitions import ContestType diff --git a/generalresearch/models/thl/contest/leaderboard.py b/generalresearch/models/thl/contest/leaderboard.py index c5e0626..696cdea 100644 --- a/generalresearch/models/thl/contest/leaderboard.py +++ b/generalresearch/models/thl/contest/leaderboard.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from typing import Any, Literal, Self from pydantic import ( diff --git a/generalresearch/models/thl/contest/milestone.py b/generalresearch/models/thl/contest/milestone.py index 8b74d50..8d96fcb 100644 --- a/generalresearch/models/thl/contest/milestone.py +++ b/generalresearch/models/thl/contest/milestone.py @@ -2,7 +2,7 @@ from __future__ import annotations import logging from datetime import timedelta -from typing import Any, Literal +from typing import Any, Literal, Self from pydantic import ( BaseModel, @@ -10,7 +10,6 @@ from pydantic import ( Field, PositiveInt, ) -from typing import Self from generalresearch.models.custom_types import AwareDatetimeISO from generalresearch.models.thl.contest.contest import ( diff --git a/generalresearch/models/thl/contest/raffle.py b/generalresearch/models/thl/contest/raffle.py index b497a44..08243f4 100644 --- a/generalresearch/models/thl/contest/raffle.py +++ b/generalresearch/models/thl/contest/raffle.py @@ -3,7 +3,7 @@ from __future__ import annotations import logging import random from collections import defaultdict -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from typing import Any, Literal, Self from pydantic import ( diff --git a/generalresearch/models/thl/finance.py b/generalresearch/models/thl/finance.py index a992f78..0856825 100644 --- a/generalresearch/models/thl/finance.py +++ b/generalresearch/models/thl/finance.py @@ -1,7 +1,7 @@ from __future__ import annotations import random -from datetime import UTC, timezone +from datetime import UTC from typing import TYPE_CHECKING from uuid import uuid4 @@ -28,8 +28,7 @@ payout_example = random.randint(150, 750 * 100) adjustment_example = random.randint(-1_000, 50 * 100) if TYPE_CHECKING: - from generalresearch.managers.thl.product import ProductManager - from generalresearch.models.thl.ledger import AccountType, Direction, LedgerAccount + from generalresearch.models.thl.ledger import LedgerAccount class AdjustmentType(BaseModel): diff --git a/generalresearch/models/thl/ipinfo.py b/generalresearch/models/thl/ipinfo.py index 3f212cf..e327bae 100644 --- a/generalresearch/models/thl/ipinfo.py +++ b/generalresearch/models/thl/ipinfo.py @@ -1,7 +1,7 @@ from __future__ import annotations import ipaddress -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from typing import Any, Literal, Self from faker import Faker diff --git a/generalresearch/models/thl/ledger_example.py b/generalresearch/models/thl/ledger_example.py index 92ad83d..0291691 100644 --- a/generalresearch/models/thl/ledger_example.py +++ b/generalresearch/models/thl/ledger_example.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from typing import Any from uuid import uuid4 diff --git a/generalresearch/models/thl/offerwall/cache.py b/generalresearch/models/thl/offerwall/cache.py index 82ab36d..97546b2 100644 --- a/generalresearch/models/thl/offerwall/cache.py +++ b/generalresearch/models/thl/offerwall/cache.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from typing import Any from pydantic import BaseModel, Field diff --git a/generalresearch/models/thl/payout_format.py b/generalresearch/models/thl/payout_format.py index a7fc4fb..4d616b6 100644 --- a/generalresearch/models/thl/payout_format.py +++ b/generalresearch/models/thl/payout_format.py @@ -45,7 +45,7 @@ def format_payout_format(payout_format: str, payout_int: int) -> str: try: xform, formatstr = inside.split(":") - except ValueError as e: + except ValueError: raise ValueError( "Payout format string must contain ':' to distinguish between transformations and formatting." ) @@ -61,17 +61,17 @@ def format_payout_format(payout_format: str, payout_int: int) -> str: payout = decimal.Decimal(eval(xform, {"payout": payout_int})) - except NameError as e: + except NameError: raise ValueError("Payout format string must contain 'payout' variable.") - except ZeroDivisionError as e: + except ZeroDivisionError: raise ValueError("Cannot divide by zero.") - except TypeError as e: + except TypeError: # "{payout()*1:}" - TypeError: 'int' object is not callable raise ValueError("Invalid type reference.") - except Exception as e: - raise ValueError(f"Invalid payout transformation") + except Exception: + raise ValueError("Invalid payout transformation") formatstr = f"{{:{formatstr}}}" diff --git a/generalresearch/models/thl/product.py b/generalresearch/models/thl/product.py index 5b7ab9a..76a8e83 100644 --- a/generalresearch/models/thl/product.py +++ b/generalresearch/models/thl/product.py @@ -77,7 +77,6 @@ if TYPE_CHECKING: from generalresearch.models.thl.payout import ( BrokerageProductPayoutEvent, ) - from generalresearch.models.thl.user import User # fmt: off @@ -1088,7 +1087,6 @@ class Product(BaseModel, validate_assignment=True): from generalresearch.incite.schemas.mergers.pop_ledger import ( numerical_col_names, ) - from generalresearch.models.thl.ledger import LedgerAccount account: LedgerAccount = thl_lm.get_account_or_create_bp_wallet(product=self) assert self.id == account.reference_uuid diff --git a/generalresearch/models/thl/profiling/marketplace.py b/generalresearch/models/thl/profiling/marketplace.py index 9038cf6..ad4ce80 100644 --- a/generalresearch/models/thl/profiling/marketplace.py +++ b/generalresearch/models/thl/profiling/marketplace.py @@ -1,7 +1,7 @@ from __future__ import annotations from abc import ABC, abstractmethod -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from functools import cached_property from typing import Any diff --git a/generalresearch/models/thl/profiling/upk_question.py b/generalresearch/models/thl/profiling/upk_question.py index 08ea350..77bba6f 100644 --- a/generalresearch/models/thl/profiling/upk_question.py +++ b/generalresearch/models/thl/profiling/upk_question.py @@ -367,7 +367,7 @@ class UpkQuestion(BaseModel): self.choices is None ), f"No `choices` are allowed for type `{self.type}`" else: - assert self.choices is not None, f"`choices` must be set" + assert self.choices is not None, "`choices` must be set" return self @model_validator(mode="after") diff --git a/generalresearch/models/thl/profiling/upk_question_answer.py b/generalresearch/models/thl/profiling/upk_question_answer.py index c59d99d..d8323ad 100644 --- a/generalresearch/models/thl/profiling/upk_question_answer.py +++ b/generalresearch/models/thl/profiling/upk_question_answer.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from typing import Any, Self from uuid import uuid4 diff --git a/generalresearch/models/thl/profiling/user_question_answer.py b/generalresearch/models/thl/profiling/user_question_answer.py index a55b205..2db07b7 100644 --- a/generalresearch/models/thl/profiling/user_question_answer.py +++ b/generalresearch/models/thl/profiling/user_question_answer.py @@ -2,7 +2,7 @@ from __future__ import annotations import json from collections.abc import Iterator -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from typing import Any, Literal, Self from pydantic import ( diff --git a/generalresearch/models/thl/session.py b/generalresearch/models/thl/session.py index c8e681c..fe7194a 100644 --- a/generalresearch/models/thl/session.py +++ b/generalresearch/models/thl/session.py @@ -2,7 +2,7 @@ from __future__ import annotations import json import logging -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal from typing import TYPE_CHECKING, Annotated, Any, Self from uuid import uuid4 diff --git a/generalresearch/models/thl/survey/__init__.py b/generalresearch/models/thl/survey/__init__.py index bc3d699..6e6b475 100644 --- a/generalresearch/models/thl/survey/__init__.py +++ b/generalresearch/models/thl/survey/__init__.py @@ -111,7 +111,6 @@ class MarketplaceTask(BaseModel, ABC): """ The Condition Model for this survey class """ - pass @property @abstractmethod @@ -119,7 +118,6 @@ class MarketplaceTask(BaseModel, ABC): """ The age question ID """ - pass @property @abstractmethod @@ -129,7 +127,6 @@ class MarketplaceTask(BaseModel, ABC): """ Mapping of generic Gender to the marketplace condition for that gender """ - pass @property def marketplace_age_groups( diff --git a/generalresearch/models/thl/survey/buyer.py b/generalresearch/models/thl/survey/buyer.py index 384bab4..b888007 100644 --- a/generalresearch/models/thl/survey/buyer.py +++ b/generalresearch/models/thl/survey/buyer.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from decimal import Decimal from math import log from typing import Annotated diff --git a/generalresearch/models/thl/survey/model.py b/generalresearch/models/thl/survey/model.py index 57bcbe2..2eed8f7 100644 --- a/generalresearch/models/thl/survey/model.py +++ b/generalresearch/models/thl/survey/model.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from decimal import Decimal from typing import Annotated, Any diff --git a/generalresearch/models/thl/survey/penalty.py b/generalresearch/models/thl/survey/penalty.py index 05153fe..04f8e20 100644 --- a/generalresearch/models/thl/survey/penalty.py +++ b/generalresearch/models/thl/survey/penalty.py @@ -1,7 +1,7 @@ from __future__ import annotations import abc -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from typing import Annotated, Literal from pydantic import BaseModel, ConfigDict, Field, TypeAdapter diff --git a/generalresearch/models/thl/task_adjustment.py b/generalresearch/models/thl/task_adjustment.py index 1834898..27c47d4 100644 --- a/generalresearch/models/thl/task_adjustment.py +++ b/generalresearch/models/thl/task_adjustment.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from decimal import Decimal from uuid import uuid4 diff --git a/generalresearch/models/thl/user.py b/generalresearch/models/thl/user.py index 55bbd18..11e0d67 100644 --- a/generalresearch/models/thl/user.py +++ b/generalresearch/models/thl/user.py @@ -3,7 +3,7 @@ from __future__ import annotations import json import logging import re -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from typing import TYPE_CHECKING, Annotated, Self from uuid import UUID, uuid4 diff --git a/generalresearch/models/thl/user_iphistory.py b/generalresearch/models/thl/user_iphistory.py index 469f8ba..257c41a 100644 --- a/generalresearch/models/thl/user_iphistory.py +++ b/generalresearch/models/thl/user_iphistory.py @@ -1,7 +1,7 @@ from __future__ import annotations import ipaddress -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from typing import Self from faker import Faker @@ -204,7 +204,6 @@ class UserIPHistory(BaseModel): if res.get(x.ip): x.information = res[x.ip] - return None def collapse_ip_records(self): """ diff --git a/generalresearch/models/thl/wallet/payout.py b/generalresearch/models/thl/wallet/payout.py index cbb37fe..42530b3 100644 --- a/generalresearch/models/thl/wallet/payout.py +++ b/generalresearch/models/thl/wallet/payout.py @@ -2,7 +2,7 @@ from __future__ import annotations import json from collections.abc import Collection -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from typing import Any from uuid import uuid4 diff --git a/generalresearch/pg_helper.py b/generalresearch/pg_helper.py index b5e124a..1d5d30b 100644 --- a/generalresearch/pg_helper.py +++ b/generalresearch/pg_helper.py @@ -1,12 +1,11 @@ from __future__ import annotations -from datetime import UTC, timezone +from datetime import UTC import psycopg -from psycopg.adapt import Buffer from psycopg.rows import RowFactory, dict_row from psycopg.types.datetime import TimestampLoader -from psycopg.types.net import Address, InetLoader, Interface +from psycopg.types.net import InetLoader from psycopg.types.string import TextLoader from psycopg.types.uuid import UUIDLoader from pydantic import PostgresDsn @@ -103,10 +102,9 @@ class PostgresConfig: # This is only intended for SELECT queries assert "SELECT" in query.upper(), "Supports SELECTs only" - with self.make_connection() as conn: - with conn.cursor() as c: - c.execute(query=query, params=params) - return c.fetchall() + with self.make_connection() as conn, conn.cursor() as c: + c.execute(query=query, params=params) + return c.fetchall() def execute_write(self, query, params=None) -> int: cmd = query.lstrip().upper() diff --git a/generalresearch/schemas/survey_stats.py b/generalresearch/schemas/survey_stats.py index 0d509e1..b3acf34 100644 --- a/generalresearch/schemas/survey_stats.py +++ b/generalresearch/schemas/survey_stats.py @@ -70,7 +70,7 @@ UnitInterval = Column( SID_CHECKS = [ Check.str_length(min_value=3, max_value=67), - Check.str_matches("^[a-z]{1,2}\:[A-Za-z0-9]+"), + Check.str_matches(r"^[a-z]{1,2}\:[A-Za-z0-9]+"), Check( lambda x: len(set(x.str.split(":").str[0])) == 1, error="the sources must all be the same", diff --git a/generalresearch/sql_helper.py b/generalresearch/sql_helper.py index 2bbdbe5..ea45305 100644 --- a/generalresearch/sql_helper.py +++ b/generalresearch/sql_helper.py @@ -1,7 +1,7 @@ from __future__ import annotations import logging -from typing import Any, Optional +from typing import Any from uuid import UUID from pydantic import MariaDBDsn, MySQLDsn, PostgresDsn @@ -39,14 +39,7 @@ class SqlConnector: # I'm intentionally doing a match case here so that we'll make sure # we can NOT use this on old versions of python 😈 - if "mysql" in self.dsn.scheme: - import pymysql as engine_module - - self.engine_module = engine_module - self.cursor_class = engine_module.cursors.DictCursor - self.quote_char = "`" - - elif "maria" in self.dsn.scheme: + if "mysql" in self.dsn.scheme or "maria" in self.dsn.scheme: import pymysql as engine_module self.engine_module = engine_module @@ -199,7 +192,6 @@ class SqlHelper(SqlConnector): if cursor is None: c.connection.commit() - return None def bulk_update( self, @@ -233,7 +225,7 @@ class SqlHelper(SqlConnector): if cursor is None: c.connection.commit() - return None + return def get_or_create( self, @@ -253,7 +245,7 @@ class SqlHelper(SqlConnector): lookup_fns = ",".join( ["`" + x + "`" for x in set(lookup_dict.keys()) | {primary_key}] ) - lookup_vals = " AND ".join([f"`{fn}`=%({fn})s" for fn in lookup_dict.keys()]) + lookup_vals = " AND ".join([f"`{fn}`=%({fn})s" for fn in lookup_dict]) table_name_str = self._quote(table_name) query = f"SELECT {lookup_fns} FROM {table_name_str} WHERE {lookup_vals} LIMIT 2" if cursor is None: @@ -293,7 +285,7 @@ class SqlHelper(SqlConnector): else: c = cursor field_names = ",".join(map(self._quote, create_dict)) - vals = ",".join([f"%({fn})s" for fn in create_dict.keys()]) + vals = ",".join([f"%({fn})s" for fn in create_dict]) table_name_str = self._quote(table_name) query = f"INSERT INTO {table_name_str} ({field_names}) VALUES ({vals})" c.execute(query, create_dict) @@ -352,4 +344,3 @@ class SqlHelper(SqlConnector): if cursor is None: c.connection.commit() - return None diff --git a/generalresearch/thl_django/apps.py b/generalresearch/thl_django/apps.py index 2813947..bd87110 100644 --- a/generalresearch/thl_django/apps.py +++ b/generalresearch/thl_django/apps.py @@ -6,11 +6,11 @@ class THLSchemaConfig(AppConfig): label = "thl_django" def ready(self): - from .accounting import models # noqa: F401 # pycharm: keep - from .common import models # noqa: F401 # pycharm: keep - from .contest import models # noqa: F401 # pycharm: keep - from .event import models # noqa: F401 # pycharm: keep - from .marketplace import models # noqa: F401 # pycharm: keep - from .network import models # noqa: F401 # pycharm: keep - from .userhealth import models # noqa: F401 # pycharm: keep + from .accounting import models # pycharm: keep + from .common import models # pycharm: keep + from .contest import models # pycharm: keep + from .event import models # pycharm: keep + from .marketplace import models # pycharm: keep + from .network import models # pycharm: keep + from .userhealth import models # pycharm: keep from .userprofile import models # noqa: F401 # pycharm: keep diff --git a/generalresearch/thl_django/fields.py b/generalresearch/thl_django/fields.py index 5e40ef0..251faa5 100644 --- a/generalresearch/thl_django/fields.py +++ b/generalresearch/thl_django/fields.py @@ -1,6 +1,7 @@ -from django.db import models import ipaddress +from django.db import models + class CIDRField(models.Field): description = "PostgreSQL CIDR network" diff --git a/generalresearch/thl_django/migrations/0001_initial.py b/generalresearch/thl_django/migrations/0001_initial.py index ecae35a..cf147e3 100644 --- a/generalresearch/thl_django/migrations/0001_initial.py +++ b/generalresearch/thl_django/migrations/0001_initial.py @@ -1,7 +1,8 @@ # Generated by Django 6.0 on 2025-12-26 20:53 -import django.db.models.deletion import uuid + +import django.db.models.deletion from django.db import migrations, models diff --git a/generalresearch/thl_django/migrations/0002_surveystat_is_live_alter_surveycategory_strength_and_more.py b/generalresearch/thl_django/migrations/0002_surveystat_is_live_alter_surveycategory_strength_and_more.py index 211c48a..f767afc 100644 --- a/generalresearch/thl_django/migrations/0002_surveystat_is_live_alter_surveycategory_strength_and_more.py +++ b/generalresearch/thl_django/migrations/0002_surveystat_is_live_alter_surveycategory_strength_and_more.py @@ -1,7 +1,7 @@ # Generated by Django 6.0 on 2025-12-28 16:49 -from django.db import migrations, models from django.contrib.postgres.operations import AddIndexConcurrently +from django.db import migrations, models class Migration(migrations.Migration): diff --git a/generalresearch/thl_django/migrations/0003_remove_surveystat_surveystat_live_survey_idx_and_more.py b/generalresearch/thl_django/migrations/0003_remove_surveystat_surveystat_live_survey_idx_and_more.py index ecaf0a9..dcf9ef2 100644 --- a/generalresearch/thl_django/migrations/0003_remove_surveystat_surveystat_live_survey_idx_and_more.py +++ b/generalresearch/thl_django/migrations/0003_remove_surveystat_surveystat_live_survey_idx_and_more.py @@ -1,10 +1,10 @@ # Generated by Django 6.0 on 2025-12-29 21:22 -from django.db import migrations, models from django.contrib.postgres.operations import ( AddIndexConcurrently, RemoveIndexConcurrently, ) +from django.db import migrations, models class Migration(migrations.Migration): diff --git a/generalresearch/thl_django/migrations/0006_remove_thlsession_thl_session_status_d578b7_idx_and_more.py b/generalresearch/thl_django/migrations/0006_remove_thlsession_thl_session_status_d578b7_idx_and_more.py index e2492ab..64338c8 100644 --- a/generalresearch/thl_django/migrations/0006_remove_thlsession_thl_session_status_d578b7_idx_and_more.py +++ b/generalresearch/thl_django/migrations/0006_remove_thlsession_thl_session_status_d578b7_idx_and_more.py @@ -1,7 +1,7 @@ # Generated by Django 6.0 on 2026-01-02 17:38 -from django.db import migrations from django.contrib.postgres.operations import RemoveIndexConcurrently +from django.db import migrations class Migration(migrations.Migration): diff --git a/generalresearch/thl_django/migrations/0009_toolrun_mtrhop_portscanport_iplabel_mtr_portscan_and_more.py b/generalresearch/thl_django/migrations/0009_toolrun_mtrhop_portscanport_iplabel_mtr_portscan_and_more.py index e19a353..e8ac2c2 100644 --- a/generalresearch/thl_django/migrations/0009_toolrun_mtrhop_portscanport_iplabel_mtr_portscan_and_more.py +++ b/generalresearch/thl_django/migrations/0009_toolrun_mtrhop_portscanport_iplabel_mtr_portscan_and_more.py @@ -1,13 +1,14 @@ # Generated by Django 6.0 on 2026-03-15 20:17 +import uuid + import django.contrib.postgres.indexes import django.db.models.deletion import django.utils.timezone from django.contrib.postgres.operations import CreateExtension +from django.db import migrations, models import generalresearch.thl_django.fields -import uuid -from django.db import migrations, models class Migration(migrations.Migration): diff --git a/generalresearch/thl_django/network/models.py b/generalresearch/thl_django/network/models.py index 167af02..733c0ab 100644 --- a/generalresearch/thl_django/network/models.py +++ b/generalresearch/thl_django/network/models.py @@ -1,12 +1,11 @@ from uuid import uuid4 -from django.utils import timezone -from django.contrib.postgres.indexes import GistIndex, GinIndex +from django.contrib.postgres.indexes import GinIndex, GistIndex from django.db import models +from django.utils import timezone from generalresearch.thl_django.fields import CIDRField - ####### # ** Signals ** # ToolRun diff --git a/generalresearch/utils/enum.py b/generalresearch/utils/enum.py index e59a383..b4620b3 100644 --- a/generalresearch/utils/enum.py +++ b/generalresearch/utils/enum.py @@ -19,7 +19,7 @@ class ReprEnumMeta(EnumMeta): [f" - __{e.value}__ *({e.name})*: {descriptions[e.name]}" for e in self] ) else: - return f"\nAllowed values: \n" + "\n".join( + return "\nAllowed values: \n" + "\n".join( [f" - __{e.value}__ *({e.name})*: {descriptions[e.name]}" for e in self] ) @@ -35,7 +35,7 @@ class ReprEnumMeta(EnumMeta): [f" - __{e.name}__: {descriptions[e.name]}" for e in self] ) else: - return f"\nAllowed values: \n" + "\n".join( + return "\nAllowed values: \n" + "\n".join( [f" - __{e.name}__: {descriptions[e.name]}" for e in self] ) diff --git a/generalresearch/wall_status_codes/lucid.py b/generalresearch/wall_status_codes/lucid.py index 05b75fb..3cc1b5e 100644 --- a/generalresearch/wall_status_codes/lucid.py +++ b/generalresearch/wall_status_codes/lucid.py @@ -61,7 +61,7 @@ client_status_map: dict[str, StatusCode1] = { "35": StatusCode1.BUYER_QUALITY_FAIL, } -status_map = defaultdict(lambda: Status.FAIL, **{"s": Status.COMPLETE}) +status_map = defaultdict(lambda: Status.FAIL, s=Status.COMPLETE) status_codes_ext_map: dict[StatusCode1, list[str]] = { StatusCode1.COMPLETE: [], StatusCode1.BUYER_FAIL: ["3"], diff --git a/generalresearch/wall_status_codes/morning.py b/generalresearch/wall_status_codes/morning.py index bac9669..ffa6be2 100644 --- a/generalresearch/wall_status_codes/morning.py +++ b/generalresearch/wall_status_codes/morning.py @@ -52,7 +52,7 @@ short_code_to_status_codes_morning: dict[str, str] = { "sur_tim": "survey_timeout", "tem_ban": "temporarily_banned", } -status_map = defaultdict(lambda: Status.FAIL, **{"complete": Status.COMPLETE}) +status_map = defaultdict(lambda: Status.FAIL, complete=Status.COMPLETE) status_codes_ext_map: dict[StatusCode1, list[str]] = { StatusCode1.COMPLETE: ["complete"], diff --git a/generalresearch/wall_status_codes/pollfish.py b/generalresearch/wall_status_codes/pollfish.py index 1ae732d..a5c6e25 100644 --- a/generalresearch/wall_status_codes/pollfish.py +++ b/generalresearch/wall_status_codes/pollfish.py @@ -28,7 +28,7 @@ status_codes_map: dict[str, str] = { "su_al_ta": "survey_already_taken", "complete": "complete", } -status_map = defaultdict(lambda: Status.FAIL, **{"complete": Status.COMPLETE}) +status_map = defaultdict(lambda: Status.FAIL, complete=Status.COMPLETE) status_codes_ext_map: dict[StatusCode1, list[str]] = { StatusCode1.COMPLETE: ["complete"], StatusCode1.BUYER_FAIL: ["third_party_termination", "screenout"], diff --git a/generalresearch/wall_status_codes/precision.py b/generalresearch/wall_status_codes/precision.py index b033cf2..c7593cf 100644 --- a/generalresearch/wall_status_codes/precision.py +++ b/generalresearch/wall_status_codes/precision.py @@ -45,7 +45,7 @@ status_codes_precision: dict[str, str] = { "60": "Client Reject", "80": "Final Complete", } -status_map = defaultdict(lambda: Status.FAIL, **{"s": Status.COMPLETE}) +status_map = defaultdict(lambda: Status.FAIL, s=Status.COMPLETE) status_codes_ext_map: dict[StatusCode1, list[str]] = { StatusCode1.COMPLETE: ["10"], StatusCode1.BUYER_FAIL: ["20", "30"], diff --git a/generalresearch/wall_status_codes/repdata.py b/generalresearch/wall_status_codes/repdata.py index ad64338..24aab14 100644 --- a/generalresearch/wall_status_codes/repdata.py +++ b/generalresearch/wall_status_codes/repdata.py @@ -46,7 +46,7 @@ rd_threat_name: dict[str, str] = { "18": "MaxMind Failure", } -status_map = defaultdict(lambda: Status.FAIL, **{"complete": Status.COMPLETE}) +status_map = defaultdict(lambda: Status.FAIL, complete=Status.COMPLETE) status_code_map: dict[StatusCode1, list[str]] = { StatusCode1.COMPLETE: ["1000"], StatusCode1.BUYER_FAIL: ["2000", "4000"], diff --git a/test_utils/conftest.py b/test_utils/conftest.py index ffe458c..03e9305 100644 --- a/test_utils/conftest.py +++ b/test_utils/conftest.py @@ -5,9 +5,8 @@ import shutil import stat import subprocess import sys -import tempfile from collections.abc import Callable, Generator -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from os.path import join as pjoin from pathlib import Path from uuid import uuid4 diff --git a/test_utils/grliq/conftest.py b/test_utils/grliq/conftest.py index 7665b52..9a3bc56 100644 --- a/test_utils/grliq/conftest.py +++ b/test_utils/grliq/conftest.py @@ -1,7 +1,7 @@ from __future__ import annotations from collections.abc import Callable -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from uuid import uuid4 import pytest diff --git a/test_utils/incite/conftest.py b/test_utils/incite/conftest.py index 87ea7ae..2968d18 100644 --- a/test_utils/incite/conftest.py +++ b/test_utils/incite/conftest.py @@ -1,7 +1,7 @@ from __future__ import annotations from collections.abc import Callable -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from os.path import join as pjoin from pathlib import Path from random import choice as randchoice diff --git a/test_utils/managers/conftest.py b/test_utils/managers/conftest.py index 4dacb29..2a9ea00 100644 --- a/test_utils/managers/conftest.py +++ b/test_utils/managers/conftest.py @@ -20,14 +20,12 @@ from generalresearch.managers.thl.ipinfo import ( IPGeonameManager, IPInformationManager, ) -from generalresearch.managers.thl.profiling.uqa import UQAManager from generalresearch.managers.thl.userhealth import ( AuditLogManager, IPRecordManager, UserIpHistoryManager, ) from generalresearch.models import Source -from generalresearch.models.thl.user import User from generalresearch.pg_helper import PostgresConfig from generalresearch.redis_helper import RedisConfig from generalresearch.sql_helper import SqlHelper diff --git a/test_utils/models/conftest.py b/test_utils/models/conftest.py index 3a9e45c..5570b40 100644 --- a/test_utils/models/conftest.py +++ b/test_utils/models/conftest.py @@ -1,7 +1,7 @@ from __future__ import annotations from collections.abc import Callable -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal from random import choice as randchoice from random import randint @@ -19,17 +19,15 @@ from generalresearch.models.thl.definitions import ( ) from generalresearch.models.thl.survey.model import Buyer, Survey from generalresearch.pg_helper import PostgresConfig -from generalresearch.redis_helper import RedisConfig if TYPE_CHECKING: from generalresearch.currency import USDCent - from generalresearch.managers.gr.authentication import GRTokenManager, GRUserManager from generalresearch.managers.gr.business import ( BusinessAddressManager, BusinessBankAccountManager, BusinessManager, ) - from generalresearch.managers.gr.team import MembershipManager, TeamManager + from generalresearch.managers.gr.team import TeamManager from generalresearch.managers.thl.buyer import BuyerManager from generalresearch.managers.thl.ipinfo import ( IPGeonameManager, @@ -45,13 +43,12 @@ if TYPE_CHECKING: from generalresearch.managers.thl.user_manager.user_manager import UserManager from generalresearch.managers.thl.userhealth import AuditLogManager, IPRecordManager from generalresearch.managers.thl.wall import WallManager - from generalresearch.models.gr.authentication import GRToken, GRUser from generalresearch.models.gr.business import ( Business, BusinessAddress, BusinessBankAccount, ) - from generalresearch.models.gr.team import Membership, Team + from generalresearch.models.gr.team import Team from generalresearch.models.thl.ipinfo import IPGeoname, IPInformation from generalresearch.models.thl.payout import ( BrokerageProductPayoutEvent, diff --git a/test_utils/models/contest/conftest.py b/test_utils/models/contest/conftest.py index 0060946..b14e126 100644 --- a/test_utils/models/contest/conftest.py +++ b/test_utils/models/contest/conftest.py @@ -1,7 +1,7 @@ from __future__ import annotations from collections.abc import Callable -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from decimal import Decimal from uuid import uuid4 diff --git a/test_utils/models/network/conftest.py b/test_utils/models/network/conftest.py index cabd8dc..6ba37a3 100644 --- a/test_utils/models/network/conftest.py +++ b/test_utils/models/network/conftest.py @@ -1,5 +1,5 @@ import os -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from uuid import uuid4 import pytest diff --git a/test_utils/models/thl/conftest.py b/test_utils/models/thl/conftest.py index a2adcce..5b21c6b 100644 --- a/test_utils/models/thl/conftest.py +++ b/test_utils/models/thl/conftest.py @@ -1,7 +1,7 @@ from __future__ import annotations from collections.abc import Callable -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from decimal import ROUND_DOWN, Decimal from random import choice as rand_choice from random import choice as rchoice diff --git a/test_utils/spectrum/conftest.py b/test_utils/spectrum/conftest.py index 9c067d3..eb2e289 100644 --- a/test_utils/spectrum/conftest.py +++ b/test_utils/spectrum/conftest.py @@ -1,6 +1,6 @@ import logging import time -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from typing import TYPE_CHECKING import pytest diff --git a/tests/incite/collections/test_df_collection_base.py b/tests/incite/collections/test_df_collection_base.py index c3c64e4..176bf4b 100644 --- a/tests/incite/collections/test_df_collection_base.py +++ b/tests/incite/collections/test_df_collection_base.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from typing import TYPE_CHECKING import pandas as pd @@ -9,7 +9,6 @@ from generalresearch.incite.collections import ( DFCollection, DFCollectionType, ) -from test_utils.incite.conftest import mnt_filepath if TYPE_CHECKING: from generalresearch.incite.base import GRLDatasets diff --git a/tests/incite/collections/test_df_collection_item_base.py b/tests/incite/collections/test_df_collection_item_base.py index 8cf719d..0218f30 100644 --- a/tests/incite/collections/test_df_collection_item_base.py +++ b/tests/incite/collections/test_df_collection_item_base.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from typing import TYPE_CHECKING import pytest diff --git a/tests/incite/collections/test_df_collection_item_thl_web.py b/tests/incite/collections/test_df_collection_item_thl_web.py index 062171d..3d70e56 100644 --- a/tests/incite/collections/test_df_collection_item_thl_web.py +++ b/tests/incite/collections/test_df_collection_item_thl_web.py @@ -1,7 +1,7 @@ from __future__ import annotations from collections.abc import Callable, Generator -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from itertools import product as iter_product from os.path import join as pjoin from pathlib import Path, PurePath @@ -236,7 +236,6 @@ class TestDFCollectionItemMethod: incite_item_factory, delete_df_collection, ): - from generalresearch.models.thl.user import User if df_collection.data_type in unsupported_mock_types: return @@ -280,7 +279,6 @@ class TestDFCollectionItemMethod: incite_item_factory, delete_df_collection, ): - from generalresearch.models.thl.user import User if df_collection.data_type in unsupported_mock_types: return @@ -333,7 +331,6 @@ class TestDFCollectionItemMethod: delete_df_collection, mnt_filepath, ): - from generalresearch.models.thl.user import User if df_collection.data_type != DFCollectionType.LEDGER: return @@ -385,7 +382,6 @@ class TestDFCollectionItemMethod: delete_df_collection, mnt_filepath, ): - from generalresearch.models.thl.user import User if df_collection.data_type in unsupported_mock_types: return @@ -766,7 +762,6 @@ class TestDFCollectionItemFunctionalTest: delete_df_collection, mnt_filepath: GRLDatasets, ): - from generalresearch.models.thl.user import User if df_collection.data_type in unsupported_mock_types: return @@ -828,7 +823,6 @@ class TestDFCollectionItemFunctionalTest: import pyarrow.parquet as pq - from generalresearch.models.thl.user import User if df_collection.data_type in unsupported_mock_types: return @@ -864,7 +858,6 @@ class TestDFCollectionItemFunctionalTest: delete_df_collection, mnt_filepath: GRLDatasets, ): - from generalresearch.models.thl.user import User delete_df_collection(coll=df_collection) df_collection._client = client_no_amm @@ -920,7 +913,6 @@ class TestDFCollectionItemFunctionalTest: """Don't allow creating an archive for data that will likely be overwritten or updated """ - from generalresearch.models.thl.user import User if df_collection.data_type in unsupported_mock_types: return @@ -987,7 +979,6 @@ class TestDFCollectionItemFunctionalTest: duration: timedelta, mnt_filepath, ): - from generalresearch.models.thl.user import User delete_df_collection(coll=df_collection) df_collection._client = client_no_amm diff --git a/tests/incite/collections/test_df_collection_thl_marketplaces.py b/tests/incite/collections/test_df_collection_thl_marketplaces.py index 2597d38..d2d3ce4 100644 --- a/tests/incite/collections/test_df_collection_thl_marketplaces.py +++ b/tests/incite/collections/test_df_collection_thl_marketplaces.py @@ -1,6 +1,5 @@ -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from itertools import product -from typing import TYPE_CHECKING import pytest from pandera.pandas import Column, DataFrameSchema, Index @@ -12,10 +11,6 @@ from generalresearch.incite.collections.thl_marketplaces import ( SagoSurveyHistoryCollection, SpectrumSurveyTimeseriesCollection, ) -from test_utils.incite.conftest import mnt_filepath - -if TYPE_CHECKING: - from generalresearch.incite.base import GRLDatasets def combo_object(): diff --git a/tests/incite/collections/test_df_collection_thl_web.py b/tests/incite/collections/test_df_collection_thl_web.py index 2cb0ba0..bcdeb83 100644 --- a/tests/incite/collections/test_df_collection_thl_web.py +++ b/tests/incite/collections/test_df_collection_thl_web.py @@ -13,9 +13,7 @@ from pandera.pandas import DataFrameSchema from generalresearch.incite.collections import DFCollection, DFCollectionType if TYPE_CHECKING: - from generalresearch.incite.base import GRLDatasets from generalresearch.incite.collections import ( - DFCollectionItem, DFCollectionType, ) diff --git a/tests/incite/mergers/foundations/test_enriched_session.py b/tests/incite/mergers/foundations/test_enriched_session.py index a0ae01e..ba11725 100644 --- a/tests/incite/mergers/foundations/test_enriched_session.py +++ b/tests/incite/mergers/foundations/test_enriched_session.py @@ -1,7 +1,6 @@ -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal from itertools import product -from typing import Optional import dask.dataframe as dd import pandas as pd @@ -11,10 +10,6 @@ from generalresearch.incite.schemas.admin_responses import ( AdminPOPSessionSchema, ) from generalresearch.pg_helper import PostgresConfig -from test_utils.incite.collections.conftest import ( - session_collection, - wall_collection, -) @pytest.mark.parametrize( diff --git a/tests/incite/mergers/foundations/test_enriched_task_adjust.py b/tests/incite/mergers/foundations/test_enriched_task_adjust.py index 96c214f..8c3a647 100644 --- a/tests/incite/mergers/foundations/test_enriched_task_adjust.py +++ b/tests/incite/mergers/foundations/test_enriched_task_adjust.py @@ -5,13 +5,6 @@ import dask.dataframe as dd import pandas as pd import pytest -from test_utils.incite.collections.conftest import ( - wall_collection, - task_adj_collection, - session_collection, -) -from test_utils.incite.mergers.conftest import enriched_wall_merge - @pytest.mark.parametrize( argnames="offset, duration,", diff --git a/tests/incite/mergers/foundations/test_enriched_wall.py b/tests/incite/mergers/foundations/test_enriched_wall.py index b421df8..0e28bce 100644 --- a/tests/incite/mergers/foundations/test_enriched_wall.py +++ b/tests/incite/mergers/foundations/test_enriched_wall.py @@ -1,34 +1,15 @@ -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal from itertools import product as iter_product -from typing import Optional import dask.dataframe as dd import pandas as pd import pytest # noinspection PyUnresolvedReferences -from distributed.utils_test import ( - cleanup, - client, - client_no_amm, - cluster_fixture, - gen_cluster, - loop, - loop_in_thread, -) - from generalresearch.incite.mergers.foundations.enriched_wall import ( EnrichedWallMergeItem, ) -from test_utils.incite.collections.conftest import ( - session_collection, - wall_collection, -) -from test_utils.incite.conftest import incite_item_factory -from test_utils.incite.mergers.conftest import ( - enriched_wall_merge, -) @pytest.mark.parametrize( @@ -118,7 +99,7 @@ class TestEnrichedWall: try: modified_time1 = path.stat().st_mtime - except (Exception,): + except Exception: modified_time1 = 0 item.build( diff --git a/tests/incite/mergers/foundations/test_user_id_product.py b/tests/incite/mergers/foundations/test_user_id_product.py index a696b45..10802e5 100644 --- a/tests/incite/mergers/foundations/test_user_id_product.py +++ b/tests/incite/mergers/foundations/test_user_id_product.py @@ -1,24 +1,13 @@ -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from itertools import product import pandas as pd import pytest # noinspection PyUnresolvedReferences -from distributed.utils_test import ( - cleanup, - client, - client_no_amm, - cluster_fixture, - gen_cluster, - loop, - loop_in_thread, -) - from generalresearch.incite.mergers.foundations.user_id_product import ( UserIdProductMergeItem, ) -from test_utils.incite.mergers.conftest import user_id_product_merge @pytest.mark.parametrize( @@ -51,7 +40,7 @@ class TestUserIDProduct: try: modified_time1 = path.stat().st_mtime - except (Exception,): + except Exception: modified_time1 = 0 user_id_product_merge.build(client=client_no_amm, user_coll=user_collection) diff --git a/tests/incite/mergers/test_merge_collection.py b/tests/incite/mergers/test_merge_collection.py index 77fa8c7..15fa4db 100644 --- a/tests/incite/mergers/test_merge_collection.py +++ b/tests/incite/mergers/test_merge_collection.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from itertools import product import pandas as pd @@ -9,7 +9,6 @@ from generalresearch.incite.mergers import ( MergeCollection, MergeType, ) -from test_utils.incite.conftest import mnt_filepath merge_types = list(e for e in MergeType if e != MergeType.TEST) diff --git a/tests/incite/mergers/test_merge_collection_item.py b/tests/incite/mergers/test_merge_collection_item.py index 96f8789..3d0b644 100644 --- a/tests/incite/mergers/test_merge_collection_item.py +++ b/tests/incite/mergers/test_merge_collection_item.py @@ -1,17 +1,10 @@ -from datetime import datetime, timezone, timedelta +from datetime import timedelta from itertools import product from pathlib import PurePath import pytest from generalresearch.incite.mergers import MergeCollectionItem, MergeType -from generalresearch.incite.mergers.foundations.enriched_session import ( - EnrichedSessionMerge, -) -from generalresearch.incite.mergers.foundations.enriched_wall import ( - EnrichedWallMerge, -) -from test_utils.incite.mergers.conftest import merge_collection @pytest.mark.parametrize( diff --git a/tests/incite/mergers/test_pop_ledger.py b/tests/incite/mergers/test_pop_ledger.py index 7583faf..dc01179 100644 --- a/tests/incite/mergers/test_pop_ledger.py +++ b/tests/incite/mergers/test_pop_ledger.py @@ -1,18 +1,12 @@ -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from itertools import product as iter_product -from typing import Optional import pandas as pd import pytest -from distributed.utils_test import client_no_amm from generalresearch.incite.schemas.mergers.pop_ledger import ( numerical_col_names, ) -from test_utils.incite.collections.conftest import ledger_collection -from test_utils.incite.conftest import incite_item_factory, mnt_filepath -from test_utils.incite.mergers.conftest import pop_ledger_merge -from test_utils.managers.ledger.conftest import create_main_accounts @pytest.mark.parametrize( diff --git a/tests/incite/mergers/test_ym_survey_merge.py b/tests/incite/mergers/test_ym_survey_merge.py index 9107f21..850df8a 100644 --- a/tests/incite/mergers/test_ym_survey_merge.py +++ b/tests/incite/mergers/test_ym_survey_merge.py @@ -1,25 +1,10 @@ -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from itertools import product import pandas as pd import pytest # noinspection PyUnresolvedReferences -from distributed.utils_test import ( - cleanup, - client, - client_no_amm, - cluster_fixture, - gen_cluster, - loop, - loop_in_thread, -) - -from test_utils.incite.collections.conftest import session_collection, wall_collection -from test_utils.incite.mergers.conftest import ( - enriched_session_merge, - ym_survey_wall_merge, -) @pytest.mark.parametrize( diff --git a/tests/incite/test_collection_base.py b/tests/incite/test_collection_base.py index 5a63019..7e1577a 100644 --- a/tests/incite/test_collection_base.py +++ b/tests/incite/test_collection_base.py @@ -10,7 +10,6 @@ import pytest from _pytest._code.code import ExceptionInfo from generalresearch.incite.base import CollectionBase -from test_utils.incite.conftest import mnt_filepath AGO_15min = (datetime.now(tz=UTC) - timedelta(minutes=15)).replace(microsecond=0) AGO_1HR = (datetime.now(tz=UTC) - timedelta(hours=1)).replace(microsecond=0) @@ -108,7 +107,7 @@ class TestCollectionBase: with pytest.raises(expected_exception=ValueError) as cm: cm: ExceptionInfo - CollectionBase(offset=f"59sec", archive_path=mnt_filepath.data_src) + CollectionBase(offset="59sec", archive_path=mnt_filepath.data_src) assert "Must be equal to, or longer than 1 min" in str(cm.value) with pytest.raises(expected_exception=ValueError) as cm: diff --git a/tests/incite/test_collection_base_item.py b/tests/incite/test_collection_base_item.py index 3f4d023..7a0a581 100644 --- a/tests/incite/test_collection_base_item.py +++ b/tests/incite/test_collection_base_item.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from os.path import join as pjoin from pathlib import Path from uuid import uuid4 diff --git a/tests/incite/test_grl_flow.py b/tests/incite/test_grl_flow.py index c632f9a..6aea182 100644 --- a/tests/incite/test_grl_flow.py +++ b/tests/incite/test_grl_flow.py @@ -1,15 +1,16 @@ class TestGRLFlow: def test_init(self, mnt_filepath, thl_web_rr): + from generalresearch.incite.collections.thl_web import ( + LedgerDFCollection, + TaskAdjustmentDFCollection, + ) from generalresearch.incite.defaults import ( ledger_df_collection, task_df_collection, - pop_ledger as plm, ) - - from generalresearch.incite.collections.thl_web import ( - LedgerDFCollection, - TaskAdjustmentDFCollection, + from generalresearch.incite.defaults import ( + pop_ledger as plm, ) from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge diff --git a/tests/incite/test_interval_idx.py b/tests/incite/test_interval_idx.py index ea2bced..3034c21 100644 --- a/tests/incite/test_interval_idx.py +++ b/tests/incite/test_interval_idx.py @@ -1,5 +1,6 @@ +from datetime import datetime + import pandas as pd -from datetime import datetime, timezone, timedelta class TestIntervalIndex: diff --git a/tests/managers/gr/test_authentication.py b/tests/managers/gr/test_authentication.py index 53b6931..b9f43a6 100644 --- a/tests/managers/gr/test_authentication.py +++ b/tests/managers/gr/test_authentication.py @@ -1,11 +1,9 @@ import logging -from random import randint from uuid import uuid4 import pytest from generalresearch.models.gr.authentication import GRUser -from test_utils.models.conftest import gr_user SSO_ISSUER = "" @@ -13,7 +11,6 @@ SSO_ISSUER = "" class TestGRUserManager: def test_create(self, gr_um): - from generalresearch.models.gr.authentication import GRUser user: GRUser = gr_um.create_dummy() instance = gr_um.get_by_id(user.id) diff --git a/tests/managers/gr/test_business.py b/tests/managers/gr/test_business.py index 7eb77f8..74a5450 100644 --- a/tests/managers/gr/test_business.py +++ b/tests/managers/gr/test_business.py @@ -2,8 +2,6 @@ from uuid import uuid4 import pytest -from test_utils.models.conftest import business - class TestBusinessBankAccountManager: @@ -12,8 +10,8 @@ class TestBusinessBankAccountManager: def test_create(self, business, business_bank_account_manager): from generalresearch.models.gr.business import ( - TransferMethod, BusinessBankAccount, + TransferMethod, ) instance = business_bank_account_manager.create( diff --git a/tests/managers/gr/test_team.py b/tests/managers/gr/test_team.py index 9215da4..0918ab8 100644 --- a/tests/managers/gr/test_team.py +++ b/tests/managers/gr/test_team.py @@ -1,7 +1,5 @@ from uuid import uuid4 -from test_utils.models.conftest import team - class TestMembershipManager: diff --git a/tests/managers/leaderboard.py b/tests/managers/leaderboard.py index 149bdbb..7773030 100644 --- a/tests/managers/leaderboard.py +++ b/tests/managers/leaderboard.py @@ -1,7 +1,7 @@ import os import time import zoneinfo -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from decimal import Decimal from uuid import uuid4 diff --git a/tests/managers/network/test_label.py b/tests/managers/network/test_label.py index 5b9a790..bfc7518 100644 --- a/tests/managers/network/test_label.py +++ b/tests/managers/network/test_label.py @@ -9,8 +9,8 @@ from generalresearch.managers.network.label import IPLabelManager from generalresearch.models.network.label import ( IPLabel, IPLabelKind, - IPLabelSource, IPLabelMetadata, + IPLabelSource, ) from generalresearch.models.thl.ipinfo import normalize_ip diff --git a/tests/managers/test_events.py b/tests/managers/test_events.py index 6941c00..cc9a1bf 100644 --- a/tests/managers/test_events.py +++ b/tests/managers/test_events.py @@ -1,11 +1,10 @@ import math import random import time -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal from functools import partial from math import floor -from typing import Optional from uuid import uuid4 import pytest @@ -384,7 +383,7 @@ class TestTaskStatsManager: "task_created_count_last_24h": AggregateBySource(total=0), } event_manager.set_source_task_stats( - source=Source.TESTING, live_task_count=0, live_tasks_max_payout=Decimal("0") + source=Source.TESTING, live_task_count=0, live_tasks_max_payout=Decimal(0) ) assert event_manager.get_task_stats_raw() == { "live_task_count": AggregateBySource( @@ -400,7 +399,7 @@ class TestTaskStatsManager: event_manager.set_source_task_stats( source=Source.TESTING, live_task_count=0, - live_tasks_max_payout=Decimal("0"), + live_tasks_max_payout=Decimal(0), created_count=10, ) res = event_manager.get_task_stats_raw() @@ -414,7 +413,7 @@ class TestTaskStatsManager: event_manager.set_source_task_stats( source=Source.TESTING, live_task_count=0, - live_tasks_max_payout=Decimal("0"), + live_tasks_max_payout=Decimal(0), created_count=10, ) res = event_manager.get_task_stats_raw() @@ -428,7 +427,7 @@ class TestTaskStatsManager: event_manager.set_source_task_stats( source=Source.TESTING2, live_task_count=0, - live_tasks_max_payout=Decimal("0"), + live_tasks_max_payout=Decimal(0), created_count=1, ) res = event_manager.get_task_stats_raw() @@ -481,7 +480,7 @@ class TestChannelsSubscriptions: wall = Wall( req_survey_id="a", - req_cpi=Decimal("1"), + req_cpi=Decimal(1), source=Source.TESTING, session_id=session.id, user_id=user.user_id, @@ -497,7 +496,7 @@ class TestChannelsSubscriptions: status=Status.COMPLETE, status_code_1=StatusCode1.COMPLETE, finished=datetime.now(tz=UTC), - cpi=Decimal("1"), + cpi=Decimal(1), ) event_manager.handle_task_finish(wall, session, user) msg = event_subscriber.get_next_message() diff --git a/tests/managers/test_userpid.py b/tests/managers/test_userpid.py index 4a3f699..36c2de9 100644 --- a/tests/managers/test_userpid.py +++ b/tests/managers/test_userpid.py @@ -1,11 +1,10 @@ import pytest from pydantic import MySQLDsn -from generalresearch.managers.marketplace.user_pid import UserPidMultiManager -from generalresearch.sql_helper import SqlHelper from generalresearch.managers.cint.user_pid import CintUserPidManager from generalresearch.managers.dynata.user_pid import DynataUserPidManager from generalresearch.managers.innovate.user_pid import InnovateUserPidManager +from generalresearch.managers.marketplace.user_pid import UserPidMultiManager from generalresearch.managers.morning.user_pid import MorningUserPidManager # from generalresearch.managers.precision import PrecisionUserPidManager @@ -13,6 +12,7 @@ from generalresearch.managers.prodege.user_pid import ProdegeUserPidManager from generalresearch.managers.repdata.user_pid import RepdataUserPidManager from generalresearch.managers.sago.user_pid import SagoUserPidManager from generalresearch.managers.spectrum.user_pid import SpectrumUserPidManager +from generalresearch.sql_helper import SqlHelper dsn = "" diff --git a/tests/managers/thl/test_contest/test_leaderboard.py b/tests/managers/thl/test_contest/test_leaderboard.py index 1a52f83..7adea9c 100644 --- a/tests/managers/thl/test_contest/test_leaderboard.py +++ b/tests/managers/thl/test_contest/test_leaderboard.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from zoneinfo import ZoneInfo from generalresearch.currency import USDCent @@ -12,12 +12,6 @@ from generalresearch.models.thl.contest.leaderboard import ( ) from generalresearch.models.thl.product import Product from generalresearch.models.thl.user import User -from test_utils.managers.contest.conftest import ( - leaderboard_contest_create as contest_create, -) -from test_utils.managers.contest.conftest import ( - leaderboard_contest_in_db as contest_in_db, -) class TestLeaderboardContestCRUD: diff --git a/tests/managers/thl/test_contest/test_milestone.py b/tests/managers/thl/test_contest/test_milestone.py index 66c5dc4..ed0bbb5 100644 --- a/tests/managers/thl/test_contest/test_milestone.py +++ b/tests/managers/thl/test_contest/test_milestone.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from generalresearch.models.thl.contest.definitions import ( ContestEndReason, @@ -12,18 +12,6 @@ from generalresearch.models.thl.contest.milestone import ( ) from generalresearch.models.thl.product import Product from generalresearch.models.thl.user import User -from test_utils.managers.contest.conftest import ( - milestone_contest as contest, -) -from test_utils.managers.contest.conftest import ( - milestone_contest_create as contest_create, -) -from test_utils.managers.contest.conftest import ( - milestone_contest_factory as contest_factory, -) -from test_utils.managers.contest.conftest import ( - milestone_contest_in_db as contest_in_db, -) class TestMilestoneContest: diff --git a/tests/managers/thl/test_contest/test_raffle.py b/tests/managers/thl/test_contest/test_raffle.py index 5804ea3..736a5e9 100644 --- a/tests/managers/thl/test_contest/test_raffle.py +++ b/tests/managers/thl/test_contest/test_raffle.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime import pytest from pydantic import ValidationError @@ -28,18 +28,6 @@ from generalresearch.models.thl.contest.raffle import ( ) from generalresearch.models.thl.product import Product from generalresearch.models.thl.user import User -from test_utils.managers.contest.conftest import ( - raffle_contest as contest, -) -from test_utils.managers.contest.conftest import ( - raffle_contest_create as contest_create, -) -from test_utils.managers.contest.conftest import ( - raffle_contest_factory as contest_factory, -) -from test_utils.managers.contest.conftest import ( - raffle_contest_in_db as contest_in_db, -) class TestRaffleContest: diff --git a/tests/managers/thl/test_harmonized_uqa.py b/tests/managers/thl/test_harmonized_uqa.py index 3b6df48..81ac080 100644 --- a/tests/managers/thl/test_harmonized_uqa.py +++ b/tests/managers/thl/test_harmonized_uqa.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime import pytest diff --git a/tests/managers/thl/test_ipinfo.py b/tests/managers/thl/test_ipinfo.py index 847b00c..61d4d19 100644 --- a/tests/managers/thl/test_ipinfo.py +++ b/tests/managers/thl/test_ipinfo.py @@ -1,9 +1,9 @@ import faker from generalresearch.managers.thl.ipinfo import ( + GeoIpInfoManager, IPGeonameManager, IPInformationManager, - GeoIpInfoManager, ) from generalresearch.models.thl.ipinfo import IPGeoname, IPInformation diff --git a/tests/managers/thl/test_ledger/test_lm_accounts.py b/tests/managers/thl/test_ledger/test_lm_accounts.py index faef5fb..11b2835 100644 --- a/tests/managers/thl/test_ledger/test_lm_accounts.py +++ b/tests/managers/thl/test_ledger/test_lm_accounts.py @@ -1,4 +1,3 @@ -from collections.abc import Callable from itertools import product as iproduct from random import randint from typing import TYPE_CHECKING @@ -22,7 +21,6 @@ from generalresearch.models.thl.ledger import ( if TYPE_CHECKING: from pydantic import PositiveInt - from generalresearch.config import GRLSettings from generalresearch.currency import LedgerCurrency from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager from generalresearch.models.custom_types import AccountType, Direction, UUIDStr @@ -32,10 +30,6 @@ if TYPE_CHECKING: LedgerAccount, LedgerTransaction, ) - from generalresearch.models.thl.product import Product - from generalresearch.models.thl.session import Session - from generalresearch.models.thl.user import User - from generalresearch.models.thl.wallet import PayoutType @pytest.mark.parametrize( diff --git a/tests/managers/thl/test_ledger/test_lm_tx_locks.py b/tests/managers/thl/test_ledger/test_lm_tx_locks.py index 07c3712..020b74a 100644 --- a/tests/managers/thl/test_ledger/test_lm_tx_locks.py +++ b/tests/managers/thl/test_ledger/test_lm_tx_locks.py @@ -1,6 +1,6 @@ import logging from collections.abc import Callable -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal import pytest @@ -23,7 +23,6 @@ from generalresearch.models.thl.session import ( WallAdjustedStatus, ) from generalresearch.models.thl.user import User -from test_utils.models.conftest import product_user_wallet_no, session, user_factory logger = logging.getLogger("LedgerManager") @@ -56,7 +55,7 @@ class TestLedgerLocks: s1 = session_factory( user=user, wall_count=3, - wall_req_cpis=[Decimal("1.23"), Decimal("3.21"), Decimal("4")], + wall_req_cpis=[Decimal("1.23"), Decimal("3.21"), Decimal(4)], wall_statuses=[Status.COMPLETE, Status.COMPLETE, Status.COMPLETE], ) @@ -187,7 +186,7 @@ class TestLedgerLocks: user_id=user.user_id, source=Source.DYNATA, req_survey_id="xxx", - req_cpi=Decimal("4"), + req_cpi=Decimal(4), session_id=1, status=Status.COMPLETE, status_code_1=StatusCode1.COMPLETE, @@ -292,13 +291,7 @@ class TestLedgerLocks: status, status_code_1 = session.determine_session_status() thl_net, commission_amount, bp_pay, user_pay = session.determine_payments() session.update( - **{ - "status": status, - "status_code_1": status_code_1, - "finished": session.started + timedelta(minutes=10), - "payout": bp_pay, - "user_payout": user_pay, - } + status=status, status_code_1=status_code_1, finished=session.started + timedelta(minutes=10), payout=bp_pay, user_payout=user_pay ) print(thl_net, commission_amount, bp_pay, user_pay) @@ -336,13 +329,7 @@ class TestLedgerLocks: status, status_code_1 = session.determine_session_status() thl_net, commission_amount, bp_pay, user_pay = session.determine_payments() session.update( - **{ - "status": status, - "status_code_1": status_code_1, - "finished": session.started + timedelta(minutes=10), - "payout": bp_pay, - "user_payout": user_pay, - } + status=status, status_code_1=status_code_1, finished=session.started + timedelta(minutes=10), payout=bp_pay, user_payout=user_pay ) print(thl_net, commission_amount, bp_pay, user_pay) diff --git a/tests/managers/thl/test_ledger/test_thl_lm_accounts.py b/tests/managers/thl/test_ledger/test_thl_lm_accounts.py index 01d5fe1..8d7d828 100644 --- a/tests/managers/thl/test_ledger/test_thl_lm_accounts.py +++ b/tests/managers/thl/test_ledger/test_thl_lm_accounts.py @@ -8,9 +8,9 @@ class TestThlLedgerManagerAccounts: def test_get_account_or_create_user_wallet(self, user, thl_lm, lm): from generalresearch.currency import LedgerCurrency from generalresearch.models.thl.ledger import ( - LedgerAccount, - Direction, AccountType, + Direction, + LedgerAccount, ) account = thl_lm.get_account_or_create_user_wallet(user=user) @@ -31,9 +31,9 @@ class TestThlLedgerManagerAccounts: def test_get_account_or_create_bp_wallet(self, product, thl_lm, lm): from generalresearch.currency import LedgerCurrency from generalresearch.models.thl.ledger import ( - LedgerAccount, - Direction, AccountType, + Direction, + LedgerAccount, ) account = thl_lm.get_account_or_create_bp_wallet(product=product) @@ -54,8 +54,8 @@ class TestThlLedgerManagerAccounts: def test_get_account_or_create_bp_commission(self, product, thl_lm, lm): from generalresearch.currency import LedgerCurrency from generalresearch.models.thl.ledger import ( - Direction, AccountType, + Direction, ) account = thl_lm.get_account_or_create_bp_commission(product=product) @@ -76,8 +76,8 @@ class TestThlLedgerManagerAccounts: def test_get_account_or_create_bp_expense(self, product, expense, thl_lm, lm): from generalresearch.currency import LedgerCurrency from generalresearch.models.thl.ledger import ( - Direction, AccountType, + Direction, ) account = thl_lm.get_account_or_create_bp_expense( @@ -98,8 +98,8 @@ class TestThlLedgerManagerAccounts: def test_get_or_create_bp_pending_payout_account(self, product, thl_lm, lm): from generalresearch.currency import LedgerCurrency from generalresearch.models.thl.ledger import ( - Direction, AccountType, + Direction, ) account = thl_lm.get_or_create_bp_pending_payout_account(product=product) @@ -132,8 +132,8 @@ class TestThlLedgerManagerAccounts: self, account_cash, account_revenue_task_complete, thl_lm, lm ): from generalresearch.models.thl.ledger import ( - LedgerAccount, AccountType, + LedgerAccount, ) res = thl_lm.get_account_task_complete_revenue() @@ -155,8 +155,8 @@ class TestThlLedgerManagerAccounts: def test_get_account_cash(self, account_cash, thl_lm, lm): from generalresearch.models.thl.ledger import ( - LedgerAccount, AccountType, + LedgerAccount, ) res = thl_lm.get_account_cash() @@ -167,10 +167,10 @@ class TestThlLedgerManagerAccounts: assert res.display_name == "Operating Cash Account" def test_get_accounts(self, setup_accounts, product, user_factory, thl_lm, lm, lam): - from generalresearch.models.thl.user import User from generalresearch.managers.thl.ledger_manager.exceptions import ( LedgerAccountDoesntExistError, ) + from generalresearch.models.thl.user import User user1: User = user_factory(product=product) user2: User = user_factory(product=product) @@ -276,9 +276,9 @@ class TestLedgerAccountManager: LedgerAccountDoesntExistError, ) from generalresearch.models.thl.ledger import ( - LedgerAccount, - Direction, AccountType, + Direction, + LedgerAccount, ) u = uuid4().hex @@ -311,8 +311,8 @@ class TestLedgerAccountManager: LedgerAccountDoesntExistError, ) from generalresearch.models.thl.ledger import ( - LedgerAccount, AccountType, + LedgerAccount, ) with pytest.raises(LedgerAccountDoesntExistError): @@ -326,10 +326,10 @@ class TestLedgerAccountManager: assert user.product.uuid == account.reference_uuid def test_get_many(self, product_factory, thl_lm, lm, lam, currency): - from generalresearch.models.thl.product import Product from generalresearch.managers.thl.ledger_manager.exceptions import ( LedgerAccountDoesntExistError, ) + from generalresearch.models.thl.product import Product p1: Product = product_factory() p2: Product = product_factory() @@ -388,9 +388,9 @@ class TestLedgerAccountManager: def test_create_account(self, thl_lm, lm, lam): from generalresearch.models.thl.ledger import ( - LedgerAccount, - Direction, AccountType, + Direction, + LedgerAccount, ) u = uuid4().hex diff --git a/tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py b/tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py index 1fb9c01..cfb8f8f 100644 --- a/tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py +++ b/tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py @@ -1,5 +1,5 @@ import logging -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal from random import randint from uuid import uuid4 @@ -78,13 +78,7 @@ class TestThlLedgerManagerBPPayout: status, status_code_1 = session.determine_session_status() thl_net, commission_amount, bp_pay, user_pay = session.determine_payments() session.update( - **{ - "status": status, - "status_code_1": status_code_1, - "finished": now + timedelta(minutes=10), - "payout": bp_pay, - "user_payout": user_pay, - } + status=status, status_code_1=status_code_1, finished=now + timedelta(minutes=10), payout=bp_pay, user_payout=user_pay ) thl_lm.create_tx_bp_payment(session=session, created=wall1.started) diff --git a/tests/managers/thl/test_ledger/test_thl_lm_tx.py b/tests/managers/thl/test_ledger/test_thl_lm_tx.py index be988a1..1130621 100644 --- a/tests/managers/thl/test_ledger/test_thl_lm_tx.py +++ b/tests/managers/thl/test_ledger/test_thl_lm_tx.py @@ -1,5 +1,5 @@ import logging -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal from random import randint from uuid import uuid4 @@ -114,7 +114,7 @@ class TestThlLedgerTxManager: user_wallet_config=UserWalletConfig(amt=True, enabled=True), ) user = user_factory(product=product) - s1 = session_factory(user=user, wall_req_cpi=Decimal("1")) + s1 = session_factory(user=user, wall_req_cpi=Decimal(1)) status, status_code_1 = s1.determine_session_status() assert status == Status.COMPLETE @@ -207,9 +207,7 @@ class TestThlLedgerTxManager: # Update the finished timestamp, but nothing else. This means that # there is no financial changes needed session.update( - **{ - "finished": datetime.now(tz=UTC) + timedelta(minutes=10), - } + finished=datetime.now(tz=UTC) + timedelta(minutes=10) ) assert session.finished with caplog.at_level(logging.INFO): @@ -829,13 +827,7 @@ class TestThlLedgerTxManagerFlows: status, status_code_1 = s1.determine_session_status() thl_net, commission_amount, bp_pay, user_pay = s1.determine_payments() s1.update( - **{ - "status": status, - "status_code_1": status_code_1, - "finished": s1.started + timedelta(minutes=10), - "payout": bp_pay, - "user_payout": user_pay, - } + status=status, status_code_1=status_code_1, finished=s1.started + timedelta(minutes=10), payout=bp_pay, user_payout=user_pay ) print(thl_net, commission_amount, bp_pay, user_pay) thl_lm.create_tx_bp_payment(session=s1, created=w1.started) @@ -892,13 +884,7 @@ class TestThlLedgerTxManagerFlows: status, status_code_1 = session.determine_session_status() thl_net, commission_amount, bp_pay, user_pay = session.determine_payments() session.update( - **{ - "status": status, - "status_code_1": status_code_1, - "finished": session.started + timedelta(minutes=10), - "payout": bp_pay, - "user_payout": user_pay, - } + status=status, status_code_1=status_code_1, finished=session.started + timedelta(minutes=10), payout=bp_pay, user_payout=user_pay ) print(thl_net, commission_amount, bp_pay, user_pay) @@ -930,13 +916,7 @@ class TestThlLedgerTxManagerFlows: status, status_code_1 = session.determine_session_status() # thl_net, commission_amount, bp_pay, user_pay = session.determine_payments() session.update( - **{ - "status": status, - "status_code_1": status_code_1, - "finished": session.started + timedelta(minutes=10), - "payout": Decimal("1.53"), - "user_payout": Decimal("1.53"), - } + status=status, status_code_1=status_code_1, finished=session.started + timedelta(minutes=10), payout=Decimal("1.53"), user_payout=Decimal("1.53") ) thl_lm.create_tx_bp_payment(session=session, created=wall1.started) @@ -971,13 +951,7 @@ class TestThlLedgerTxManagerFlows: status, status_code_1 = session.determine_session_status() # thl_net, commission_amount, bp_pay, user_pay = session.determine_payments() session.update( - **{ - "status": status, - "status_code_1": status_code_1, - "finished": session.started + timedelta(minutes=10), - "payout": Decimal("0.39"), - "user_payout": Decimal("0.26"), - } + status=status, status_code_1=status_code_1, finished=session.started + timedelta(minutes=10), payout=Decimal("0.39"), user_payout=Decimal("0.26") ) # with pytest.logs(logger, level=logging.WARNING) as cm: # tx = thl_lm.create_transaction_bp_payment(session, created=wall1.started) @@ -1310,13 +1284,7 @@ class TestThlLedgerManagerAdj: status, status_code_1 = session.determine_session_status() thl_net, commission_amount, bp_pay, user_pay = session.determine_payments() session.update( - **{ - "status": status, - "status_code_1": status_code_1, - "finished": utc_hour_ago + timedelta(minutes=10), - "payout": bp_pay, - "user_payout": user_pay, - } + status=status, status_code_1=status_code_1, finished=utc_hour_ago + timedelta(minutes=10), payout=bp_pay, user_payout=user_pay ) thl_lm.create_tx_bp_payment(session, created=wall1.started) @@ -1484,13 +1452,7 @@ class TestThlLedgerManagerAdj: assert user_pay == Decimal("1.52") session.update( - **{ - "status": status, - "status_code_1": status_code_1, - "finished": now + timedelta(minutes=10), - "payout": bp_pay, - "user_payout": user_pay, - } + status=status, status_code_1=status_code_1, finished=now + timedelta(minutes=10), payout=bp_pay, user_payout=user_pay ) tx = thl_lm.create_tx_bp_adjustment(session=session, created=wall1.started) @@ -1664,13 +1626,7 @@ class TestThlLedgerManagerAdj: status, status_code_1 = session.determine_session_status() thl_net, commission_amount, bp_pay, user_pay = session.determine_payments() session.update( - **{ - "status": status, - "status_code_1": status_code_1, - "finished": utc_hour_ago + timedelta(minutes=10), - "payout": bp_pay, - "user_payout": user_pay, - } + status=status, status_code_1=status_code_1, finished=utc_hour_ago + timedelta(minutes=10), payout=bp_pay, user_payout=user_pay ) thl_lm.create_tx_bp_payment(session, created=wall1.started) diff --git a/tests/managers/thl/test_ledger/test_thl_lm_tx__user_payouts.py b/tests/managers/thl/test_ledger/test_thl_lm_tx__user_payouts.py index 9253ff0..cd6ea79 100644 --- a/tests/managers/thl/test_ledger/test_thl_lm_tx__user_payouts.py +++ b/tests/managers/thl/test_ledger/test_thl_lm_tx__user_payouts.py @@ -1,5 +1,5 @@ import logging -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal from uuid import uuid4 @@ -12,7 +12,6 @@ from generalresearch.managers.thl.ledger_manager.exceptions import ( from generalresearch.models.thl.payout import UserPayoutEvent from generalresearch.models.thl.user import User from generalresearch.models.thl.wallet import PayoutType -from test_utils.managers.ledger.conftest import create_main_accounts class TestLedgerManagerAMT: diff --git a/tests/managers/thl/test_ledger/test_thl_pem.py b/tests/managers/thl/test_ledger/test_thl_pem.py index 5fb9e7d..7a4f0c9 100644 --- a/tests/managers/thl/test_ledger/test_thl_pem.py +++ b/tests/managers/thl/test_ledger/test_thl_pem.py @@ -1,6 +1,5 @@ -import uuid from random import randint -from uuid import uuid4, UUID +from uuid import UUID, uuid4 import pytest @@ -33,7 +32,6 @@ class TestThlPayoutEventManager: thl_lm, brokerage_product_payout_event_manager, ): - from generalresearch.models.thl.payout import UserPayoutEvent N_PRODUCTS = randint(3, 10) N_PAYOUT_EVENTS = randint(3, 10) @@ -73,7 +71,6 @@ class TestThlPayoutEventManager: brokerage_product_payout_event_manager, thl_lm, ): - from generalresearch.models.thl.payout import UserPayoutEvent N_PRODUCTS = randint(3, 10) N_PAYOUT_EVENTS = randint(3, 10) @@ -119,7 +116,6 @@ class TestThlPayoutEventManager: description can't be None """ from generalresearch.models.thl.payout import ( - UserPayoutEvent, PayoutType, ) @@ -174,7 +170,6 @@ class TestThlPayoutEventManager: brokerage_product_payout_event_manager, lm, ): - from generalresearch.models.thl.payout import UserPayoutEvent delete_ledger_db() create_main_accounts() diff --git a/tests/managers/thl/test_ledger/test_user_txs.py b/tests/managers/thl/test_ledger/test_user_txs.py index b4b0437..f83641e 100644 --- a/tests/managers/thl/test_ledger/test_user_txs.py +++ b/tests/managers/thl/test_ledger/test_user_txs.py @@ -1,5 +1,5 @@ from collections.abc import Callable -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime from decimal import Decimal from typing import TYPE_CHECKING from uuid import uuid4 @@ -8,7 +8,6 @@ from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerMana from generalresearch.managers.thl.user_compensate import user_compensate from generalresearch.models.thl.definitions import ( Status, - WallAdjustedStatus, ) from generalresearch.models.thl.ledger import ( TransactionType, diff --git a/tests/managers/thl/test_ledger/test_wallet.py b/tests/managers/thl/test_ledger/test_wallet.py index a0abd7c..bb49cd8 100644 --- a/tests/managers/thl/test_ledger/test_wallet.py +++ b/tests/managers/thl/test_ledger/test_wallet.py @@ -4,10 +4,10 @@ from uuid import uuid4 import pytest from generalresearch.models.thl.product import ( - UserWalletConfig, PayoutConfig, PayoutTransformation, PayoutTransformationPercentArgs, + UserWalletConfig, ) from generalresearch.models.thl.user import User diff --git a/tests/managers/thl/test_product.py b/tests/managers/thl/test_product.py index 78d5dde..f8dd44d 100644 --- a/tests/managers/thl/test_product.py +++ b/tests/managers/thl/test_product.py @@ -5,15 +5,14 @@ import pytest from generalresearch.models import Source from generalresearch.models.thl.product import ( Product, + ProfilingConfig, SourceConfig, - UserCreateConfig, SourcesConfig, - UserHealthConfig, - ProfilingConfig, - SupplyPolicy, SupplyConfig, + SupplyPolicy, + UserCreateConfig, + UserHealthConfig, ) -from test_utils.models.conftest import product_factory class TestProductManagerGetMethods: diff --git a/tests/managers/thl/test_product_prod.py b/tests/managers/thl/test_product_prod.py index 7b4f677..f5aa78d 100644 --- a/tests/managers/thl/test_product_prod.py +++ b/tests/managers/thl/test_product_prod.py @@ -4,7 +4,6 @@ from uuid import uuid4 import pytest from generalresearch.models.thl.product import Product -from test_utils.models.conftest import product_factory logger = logging.getLogger() @@ -79,4 +78,3 @@ class TestProductManagerGetAll: products = product_manager.get_all(rand_limit=50) logger.info(f"Fetching {len(products)} product uuids") # todo: once timebucks stops spamming broken accounts, fetch more - pass diff --git a/tests/managers/thl/test_profiling/test_user_upk.py b/tests/managers/thl/test_profiling/test_user_upk.py index 491e2b1..8b995b1 100644 --- a/tests/managers/thl/test_profiling/test_user_upk.py +++ b/tests/managers/thl/test_profiling/test_user_upk.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from generalresearch.managers.thl.profiling.user_upk import UserUpkManager diff --git a/tests/managers/thl/test_session_manager.py b/tests/managers/thl/test_session_manager.py index 6bedc2b..adcbe25 100644 --- a/tests/managers/thl/test_session_manager.py +++ b/tests/managers/thl/test_session_manager.py @@ -7,11 +7,10 @@ from faker import Faker from generalresearch.models import DeviceType from generalresearch.models.legacy.bucket import Bucket from generalresearch.models.thl.definitions import ( + SessionStatusCode2, Status, StatusCode1, - SessionStatusCode2, ) -from test_utils.models.conftest import user fake = Faker() @@ -21,8 +20,8 @@ class TestSessionManager: bucket = Bucket( loi_min=timedelta(seconds=60), loi_max=timedelta(seconds=120), - user_payout_min=Decimal("1"), - user_payout_max=Decimal("2"), + user_payout_min=Decimal(1), + user_payout_max=Decimal(2), ) s1 = session_manager.create( diff --git a/tests/managers/thl/test_survey.py b/tests/managers/thl/test_survey.py index 4b4a579..37f0b66 100644 --- a/tests/managers/thl/test_survey.py +++ b/tests/managers/thl/test_survey.py @@ -1,5 +1,5 @@ import uuid -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from decimal import Decimal import pytest diff --git a/tests/managers/thl/test_survey_penalty.py b/tests/managers/thl/test_survey_penalty.py index 4c7dc08..2a3cdc2 100644 --- a/tests/managers/thl/test_survey_penalty.py +++ b/tests/managers/thl/test_survey_penalty.py @@ -1,7 +1,6 @@ import uuid import pytest -from cachetools.keys import _HashedTuple from generalresearch.models import Source from generalresearch.models.thl.survey.penalty import ( diff --git a/tests/managers/thl/test_task_adjustment.py b/tests/managers/thl/test_task_adjustment.py index 71e3535..43337b6 100644 --- a/tests/managers/thl/test_task_adjustment.py +++ b/tests/managers/thl/test_task_adjustment.py @@ -1,5 +1,5 @@ import logging -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal from random import randint diff --git a/tests/managers/thl/test_task_status.py b/tests/managers/thl/test_task_status.py index 468fd5e..93a624d 100644 --- a/tests/managers/thl/test_task_status.py +++ b/tests/managers/thl/test_task_status.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal import pytest diff --git a/tests/managers/thl/test_user_manager/test_base.py b/tests/managers/thl/test_user_manager/test_base.py index 2704490..7d83c11 100644 --- a/tests/managers/thl/test_user_manager/test_base.py +++ b/tests/managers/thl/test_user_manager/test_base.py @@ -1,5 +1,5 @@ import logging -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from random import randint from uuid import uuid4 @@ -156,7 +156,7 @@ class TestCreateUserManager: # make sure thl_user row is created res_thl_user = thl_web_rw.execute_sql_query( - query=f""" + query=""" SELECT * FROM thl_user AS u WHERE u.id = %s diff --git a/tests/managers/thl/test_user_manager/test_mysql.py b/tests/managers/thl/test_user_manager/test_mysql.py index 0313bbf..d414a13 100644 --- a/tests/managers/thl/test_user_manager/test_mysql.py +++ b/tests/managers/thl/test_user_manager/test_mysql.py @@ -1,4 +1,3 @@ -from test_utils.models.conftest import user, user_manager class TestUserManagerMysqlNew: diff --git a/tests/managers/thl/test_user_manager/test_user_fetch.py b/tests/managers/thl/test_user_manager/test_user_fetch.py index a4b3d57..7c9e012 100644 --- a/tests/managers/thl/test_user_manager/test_user_fetch.py +++ b/tests/managers/thl/test_user_manager/test_user_fetch.py @@ -3,7 +3,6 @@ from uuid import uuid4 import pytest from generalresearch.models.thl.user import User -from test_utils.models.conftest import product, user_manager, user_factory class TestUserManagerFetch: diff --git a/tests/managers/thl/test_user_manager/test_user_metadata.py b/tests/managers/thl/test_user_manager/test_user_metadata.py index 91dc16a..19b3d9f 100644 --- a/tests/managers/thl/test_user_manager/test_user_metadata.py +++ b/tests/managers/thl/test_user_manager/test_user_metadata.py @@ -3,7 +3,6 @@ from uuid import uuid4 import pytest from generalresearch.models.thl.user_profile import UserMetadata -from test_utils.models.conftest import user, user_manager, user_factory class TestUserMetadataManager: diff --git a/tests/managers/thl/test_user_streak.py b/tests/managers/thl/test_user_streak.py index ef25e2b..be0729c 100644 --- a/tests/managers/thl/test_user_streak.py +++ b/tests/managers/thl/test_user_streak.py @@ -1,5 +1,5 @@ import copy -from datetime import UTC, date, datetime, timedelta, timezone +from datetime import UTC, date, datetime, timedelta from decimal import Decimal from zoneinfo import ZoneInfo diff --git a/tests/managers/thl/test_userhealth.py b/tests/managers/thl/test_userhealth.py index e2ea4a8..be8f9c8 100644 --- a/tests/managers/thl/test_userhealth.py +++ b/tests/managers/thl/test_userhealth.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from uuid import uuid4 import faker diff --git a/tests/managers/thl/test_wall_manager.py b/tests/managers/thl/test_wall_manager.py index 46d199f..5abc648 100644 --- a/tests/managers/thl/test_wall_manager.py +++ b/tests/managers/thl/test_wall_manager.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal from uuid import uuid4 @@ -10,7 +10,6 @@ from generalresearch.models.thl.session import ( Status, StatusCode1, ) -from test_utils.models.conftest import session, user class TestWallManager: @@ -92,7 +91,7 @@ class TestWallManager: source=Source.DYNATA, buyer_id="123", req_survey_id="456", - req_cpi=Decimal("1"), + req_cpi=Decimal(1), ) assert w is not None @@ -110,7 +109,7 @@ class TestWallManager: source=Source.DYNATA, buyer_id="123", req_survey_id="456", - req_cpi=Decimal("1"), + req_cpi=Decimal(1), ) wall_manager.report( wall=w1, @@ -151,7 +150,7 @@ class TestWallManager: source=Source.DYNATA, buyer_id="123", req_survey_id="456", - req_cpi=Decimal("1"), + req_cpi=Decimal(1), ) finish_ts = utc_hour_ago + timedelta(minutes=10) @@ -190,7 +189,7 @@ class TestWallManager: source=Source.DYNATA, buyer_id="123", req_survey_id="456", - req_cpi=Decimal("1"), + req_cpi=Decimal(1), ) res = wall_manager.filter_wall_attempts(user_id=user.user_id) assert len(res) == 1 @@ -202,7 +201,7 @@ class TestWallManager: source=Source.DYNATA, buyer_id="123", req_survey_id="555", - req_cpi=Decimal("1"), + req_cpi=Decimal(1), ) res = wall_manager.filter_wall_attempts(user_id=user.user_id) assert len(res) == 2 diff --git a/tests/models/admin/test_report_request.py b/tests/models/admin/test_report_request.py index 4626ab4..5b2ff0d 100644 --- a/tests/models/admin/test_report_request.py +++ b/tests/models/admin/test_report_request.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime import pandas as pd import pytest diff --git a/tests/models/custom_types/test_aware_datetime.py b/tests/models/custom_types/test_aware_datetime.py index 043fba0..7c45710 100644 --- a/tests/models/custom_types/test_aware_datetime.py +++ b/tests/models/custom_types/test_aware_datetime.py @@ -1,7 +1,7 @@ from __future__ import annotations import logging -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime import pytest import pytz diff --git a/tests/models/custom_types/test_dsn.py b/tests/models/custom_types/test_dsn.py index 050976e..eb27526 100644 --- a/tests/models/custom_types/test_dsn.py +++ b/tests/models/custom_types/test_dsn.py @@ -1,4 +1,3 @@ -from typing import Optional from uuid import uuid4 import pytest diff --git a/tests/models/dynata/test_eligbility.py b/tests/models/dynata/test_eligbility.py index 16cad26..27de5b3 100644 --- a/tests/models/dynata/test_eligbility.py +++ b/tests/models/dynata/test_eligbility.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime class TestEligibility: diff --git a/tests/models/gr/test_authentication.py b/tests/models/gr/test_authentication.py index 51595a7..ac39e64 100644 --- a/tests/models/gr/test_authentication.py +++ b/tests/models/gr/test_authentication.py @@ -2,7 +2,7 @@ import binascii import json import os from collections.abc import Callable -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from random import randint from uuid import uuid4 diff --git a/tests/models/gr/test_business.py b/tests/models/gr/test_business.py index 716ec75..948acb3 100644 --- a/tests/models/gr/test_business.py +++ b/tests/models/gr/test_business.py @@ -1,7 +1,6 @@ import os -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal -from typing import Optional from uuid import uuid4 import pandas as pd diff --git a/tests/models/innovate/test_question.py b/tests/models/innovate/test_question.py index 330f919..b0c2964 100644 --- a/tests/models/innovate/test_question.py +++ b/tests/models/innovate/test_question.py @@ -1,15 +1,15 @@ from generalresearch.models import Source from generalresearch.models.innovate.question import ( InnovateQuestion, - InnovateQuestionType, InnovateQuestionOption, + InnovateQuestionType, ) from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestionSelectorTE, UpkQuestion, + UpkQuestionChoice, UpkQuestionSelectorMC, + UpkQuestionSelectorTE, UpkQuestionType, - UpkQuestionChoice, ) diff --git a/tests/models/legacy/test_user_question_answer_in.py b/tests/models/legacy/test_user_question_answer_in.py index 224334a..1904798 100644 --- a/tests/models/legacy/test_user_question_answer_in.py +++ b/tests/models/legacy/test_user_question_answer_in.py @@ -263,12 +263,12 @@ class TestUserQuestionAnswerIn: UserQuestionAnswerIn, ) - for qid in { + for qid in ( "2fbedb2b9f7647b09ff5e52fa119cc5e", "4030c52371b04e80b64e058d9c5b82e9", "a91cb1dea814480dba12d9b7b48696dd", "1d1e2e8380ac474b87fb4e4c569b48df", - }: + ): # This is the UserAgent question which only allows a single answer with pytest.raises(ValueError) as cm: UserQuestionAnswerIn.model_validate( diff --git a/tests/models/morning/test.py b/tests/models/morning/test.py index 222cb93..7474766 100644 --- a/tests/models/morning/test.py +++ b/tests/models/morning/test.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from generalresearch.models.morning.question import MorningQuestion diff --git a/tests/models/network/test_mtr.py b/tests/models/network/test_mtr.py index 2965300..840a773 100644 --- a/tests/models/network/test_mtr.py +++ b/tests/models/network/test_mtr.py @@ -1,7 +1,7 @@ -from generalresearch.models.network.mtr.execute import execute_mtr import faker -from generalresearch.models.network.tool_run import ToolName, ToolClass +from generalresearch.models.network.mtr.execute import execute_mtr +from generalresearch.models.network.tool_run import ToolClass, ToolName fake = faker.Faker() diff --git a/tests/models/network/test_nmap_parser.py b/tests/models/network/test_nmap_parser.py index abc83c9..7822380 100644 --- a/tests/models/network/test_nmap_parser.py +++ b/tests/models/network/test_nmap_parser.py @@ -4,6 +4,7 @@ import pytest from generalresearch.models.network.nmap.parser import parse_nmap_xml + @pytest.fixture def nmap_raw_output_2(request) -> str: fp = os.path.join(request.config.rootpath, "data/nmaprun2.xml") diff --git a/tests/models/prodege/test_survey_participation.py b/tests/models/prodege/test_survey_participation.py index 3b35d0c..e1ba9ab 100644 --- a/tests/models/prodege/test_survey_participation.py +++ b/tests/models/prodege/test_survey_participation.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta class TestProdegeParticipation: diff --git a/tests/models/spectrum/test_question.py b/tests/models/spectrum/test_question.py index 4f92961..57d260d 100644 --- a/tests/models/spectrum/test_question.py +++ b/tests/models/spectrum/test_question.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from generalresearch.models import Source from generalresearch.models.spectrum.question import ( diff --git a/tests/models/spectrum/test_survey.py b/tests/models/spectrum/test_survey.py index 5e095a3..7365c7e 100644 --- a/tests/models/spectrum/test_survey.py +++ b/tests/models/spectrum/test_survey.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from decimal import Decimal diff --git a/tests/models/spectrum/test_survey_manager.py b/tests/models/spectrum/test_survey_manager.py index 11970bf..ce26c44 100644 --- a/tests/models/spectrum/test_survey_manager.py +++ b/tests/models/spectrum/test_survey_manager.py @@ -1,6 +1,6 @@ import copy import logging -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from decimal import Decimal from pymysql import IntegrityError diff --git a/tests/models/test_finance.py b/tests/models/test_finance.py index 6dcd441..3d3ff3a 100644 --- a/tests/models/test_finance.py +++ b/tests/models/test_finance.py @@ -1,5 +1,5 @@ from collections.abc import Callable -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from itertools import product as iter_product from random import randint from uuid import uuid4 diff --git a/tests/models/thl/question/test_question_info.py b/tests/models/thl/question/test_question_info.py index 945ee7a..b619fc3 100644 --- a/tests/models/thl/question/test_question_info.py +++ b/tests/models/thl/question/test_question_info.py @@ -1,6 +1,6 @@ from generalresearch.models.thl.profiling.upk_property import ( - UpkProperty, ProfilingInfo, + UpkProperty, ) diff --git a/tests/models/thl/test_adjustments.py b/tests/models/thl/test_adjustments.py index c2c035d..30e9bce 100644 --- a/tests/models/thl/test_adjustments.py +++ b/tests/models/thl/test_adjustments.py @@ -1,5 +1,5 @@ from collections.abc import Callable -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal import pytest @@ -426,7 +426,7 @@ class TestAdjustments: s1 = session_factory( user=user, wall_count=1, - wall_req_cpi=Decimal("1"), + wall_req_cpi=Decimal(1), final_status=Status.COMPLETE, started=utc_hour_ago, ) @@ -525,7 +525,7 @@ class TestAdjustments: s1 = session_factory( user=user, wall_count=1, - wall_req_cpi=Decimal("1"), + wall_req_cpi=Decimal(1), final_status=Status.COMPLETE, started=utc_hour_ago, ) @@ -534,13 +534,7 @@ class TestAdjustments: status, status_code_1 = s1.determine_session_status() thl_net, commission_amount, bp_pay, user_pay = s1.determine_payments() s1.update( - **{ - "status": status, - "status_code_1": status_code_1, - "finished": utc_hour_ago + timedelta(minutes=10), - "payout": bp_pay, - "user_payout": user_pay, - } + status=status, status_code_1=status_code_1, finished=utc_hour_ago + timedelta(minutes=10), payout=bp_pay, user_payout=user_pay ) w1.update( adjusted_status=WallAdjustedStatus.ADJUSTED_TO_FAIL, @@ -592,13 +586,7 @@ class TestAdjustments: thl_net = Decimal(sum(w.cpi for w in s1.wall_events if w.is_visible_complete())) payout = user.product.determine_bp_payment(thl_net=thl_net) s1.update( - **{ - "status": status, - "status_code_1": status_code_1, - "finished": utc_hour_ago + timedelta(minutes=25), - "payout": payout, - "user_payout": None, - } + status=status, status_code_1=status_code_1, finished=utc_hour_ago + timedelta(minutes=25), payout=payout, user_payout=None ) # Test. Adjust first fail to complete. Now we have 2 completes. @@ -646,13 +634,7 @@ class TestAdjustments: thl_net = Decimal(sum(w.cpi for w in s1.wall_events if w.is_visible_complete())) payout = user.product.determine_bp_payment(thl_net) s1.update( - **{ - "status": status, - "status_code_1": status_code_1, - "finished": utc_hour_ago + timedelta(minutes=25), - "payout": payout, - "user_payout": None, - } + status=status, status_code_1=status_code_1, finished=utc_hour_ago + timedelta(minutes=25), payout=payout, user_payout=None ) # Test. Adjust complete to fail. Now we have 2 fails. diff --git a/tests/models/thl/test_contest/test_leaderboard_contest.py b/tests/models/thl/test_contest/test_leaderboard_contest.py index 3efcf2f..52a4bec 100644 --- a/tests/models/thl/test_contest/test_leaderboard_contest.py +++ b/tests/models/thl/test_contest/test_leaderboard_contest.py @@ -1,4 +1,4 @@ -from datetime import UTC, timezone +from datetime import UTC from uuid import uuid4 import pytest diff --git a/tests/models/thl/test_ledger.py b/tests/models/thl/test_ledger.py index 5edcc9d..7066180 100644 --- a/tests/models/thl/test_ledger.py +++ b/tests/models/thl/test_ledger.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from uuid import uuid4 import pytest diff --git a/tests/models/thl/test_payout.py b/tests/models/thl/test_payout.py index 7068a41..d0b0acc 100644 --- a/tests/models/thl/test_payout.py +++ b/tests/models/thl/test_payout.py @@ -5,14 +5,13 @@ from pydantic import ValidationError from generalresearch.currency import USDCent from generalresearch.models.gr import Team +from generalresearch.models.gr.business import Business, BusinessAddress, BusinessType from generalresearch.models.thl.payout import ( - BusinessPayoutEvent, BrokerageProductPayoutEvent, + BusinessPayoutEvent, ) from generalresearch.models.thl.wallet import PayoutType -from generalresearch.models.gr.business import Business, BusinessAddress, BusinessType - class TestBusinessPayoutEvent: diff --git a/tests/models/thl/test_product.py b/tests/models/thl/test_product.py index 78bc10a..bc95c2d 100644 --- a/tests/models/thl/test_product.py +++ b/tests/models/thl/test_product.py @@ -3,7 +3,7 @@ from __future__ import annotations import os import shutil from collections.abc import Callable -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal from uuid import uuid4 diff --git a/tests/models/thl/test_upkquestion.py b/tests/models/thl/test_upkquestion.py index d32875c..99d7871 100644 --- a/tests/models/thl/test_upkquestion.py +++ b/tests/models/thl/test_upkquestion.py @@ -201,18 +201,10 @@ class TestUpkQuestion: ) q = MorningQuestion( - **{ - "id": "gender", - "country_iso": "us", - "language_iso": "eng", - "name": "Gender", - "text": "What is your gender?", - "type": "s", - "options": [ + id="gender", country_iso="us", language_iso="eng", name="Gender", text="What is your gender?", type="s", options=[ {"id": "1", "text": "yes", "order": 1}, {"id": "2", "text": "no", "order": 2}, - ], - } + ] ) q.to_upk_question() q = MorningQuestion( diff --git a/tests/models/thl/test_user.py b/tests/models/thl/test_user.py index a4f331a..aafce68 100644 --- a/tests/models/thl/test_user.py +++ b/tests/models/thl/test_user.py @@ -197,7 +197,7 @@ class TestUserProductUserID: assert "Input should be a valid string" in str(cm.value) with pytest.raises(ValueError) as cm: - User(user_id=self.user_id, product_user_id=Decimal("0")) + User(user_id=self.user_id, product_user_id=Decimal(0)) assert "1 validation error for User" in str(cm.value) assert "Input should be a valid string" in str(cm.value) @@ -236,7 +236,7 @@ class TestUserProductUserID: def test_invalid_chars_slash(self): from generalresearch.models.thl.user import User - product_user_id = f"{self.randomword(50)}\{self.randomword(50)}" + product_user_id = rf"{self.randomword(50)}\{self.randomword(50)}" with pytest.raises(expected_exception=ValueError) as cm: User(user_id=self.user_id, product_user_id=product_user_id) assert "1 validation error for User" in str(cm.value) @@ -310,7 +310,7 @@ class TestUserUUID: assert "Input should be a valid string" in str(cm.value) with pytest.raises(ValueError) as cm: - User(user_id=self.user_id, uuid=Decimal("0")) + User(user_id=self.user_id, uuid=Decimal(0)) assert "1 validation error for User" in str(cm.value) assert "Input should be a valid string" in str(cm.value) diff --git a/tests/models/thl/test_user_iphistory.py b/tests/models/thl/test_user_iphistory.py index 0f050b0..d6ade9d 100644 --- a/tests/models/thl/test_user_iphistory.py +++ b/tests/models/thl/test_user_iphistory.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from generalresearch.models.thl.user_iphistory import ( UserIPHistory, diff --git a/tests/models/thl/test_user_streak.py b/tests/models/thl/test_user_streak.py index 0cacd3e..26c5e25 100644 --- a/tests/models/thl/test_user_streak.py +++ b/tests/models/thl/test_user_streak.py @@ -1,8 +1,8 @@ from datetime import datetime, timedelta +from zoneinfo import ZoneInfo import pytest from pydantic import ValidationError -from zoneinfo import ZoneInfo from generalresearch.models.thl.user_streak import ( StreakFulfillment, diff --git a/tests/models/thl/test_wall.py b/tests/models/thl/test_wall.py index 9e9483b..88914ac 100644 --- a/tests/models/thl/test_wall.py +++ b/tests/models/thl/test_wall.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal from uuid import uuid4 diff --git a/tests/models/thl/test_wall_session.py b/tests/models/thl/test_wall_session.py index 10f3cba..b39ad31 100644 --- a/tests/models/thl/test_wall_session.py +++ b/tests/models/thl/test_wall_session.py @@ -1,4 +1,4 @@ -from datetime import UTC, datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from decimal import Decimal import pytest diff --git a/tests/sql_helper.py b/tests/sql_helper.py index c4cc2ca..8ab7bdb 100644 --- a/tests/sql_helper.py +++ b/tests/sql_helper.py @@ -19,7 +19,7 @@ class TestSqlHelper: def test_scheme(self): from generalresearch.sql_helper import SqlHelper - dsn = MySQLDsn(f"mysql://root@localhost/test") + dsn = MySQLDsn("mysql://root@localhost/test") instance = SqlHelper(dsn=dsn) assert instance.is_mysql() @@ -30,7 +30,7 @@ class TestSqlHelper: # self.assertTrue(instance.is_postgresql()) with pytest.raises(ValidationError): - SqlHelper(dsn=MariaDBDsn(f"maria://root@localhost/test")) + SqlHelper(dsn=MariaDBDsn("maria://root@localhost/test")) def test_row_decode(self): from generalresearch.sql_helper import decode_uuids -- cgit v1.2.3 From 3b4059135be47f7752a08e4277a85f9e57ceaa9d Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Tue, 25 Aug 2026 11:09:52 -0700 Subject: Ruff typing from this morning. WIP --- test_utils/models/ledger/conftest.py | 117 +-- .../incite/collections/test_df_collection_base.py | 2 +- .../collections/test_df_collection_item_base.py | 2 +- .../collections/test_df_collection_item_thl_web.py | 61 +- .../mergers/foundations/test_enriched_session.py | 18 +- .../foundations/test_enriched_task_adjust.py | 12 +- .../mergers/foundations/test_enriched_wall.py | 36 +- tests/incite/mergers/test_pop_ledger.py | 34 +- tests/incite/mergers/test_ym_survey_merge.py | 12 +- tests/managers/gr/test_business.py | 12 +- tests/managers/gr/test_team.py | 4 +- tests/managers/leaderboard.py | 2 +- tests/managers/test_events.py | 8 +- tests/managers/test_lucid.py | 4 +- .../managers/thl/test_contest/test_leaderboard.py | 49 +- tests/managers/thl/test_contest/test_milestone.py | 79 +- tests/managers/thl/test_contest/test_raffle.py | 146 ++- tests/managers/thl/test_ipinfo.py | 11 +- tests/managers/thl/test_ledger/test_lm_accounts.py | 23 +- tests/managers/thl/test_ledger/test_lm_tx.py | 142 ++- .../managers/thl/test_ledger/test_lm_tx_entries.py | 24 +- tests/managers/thl/test_ledger/test_lm_tx_locks.py | 208 ++-- .../thl/test_ledger/test_lm_tx_metadata.py | 40 +- .../thl/test_ledger/test_thl_lm_accounts.py | 310 +++--- .../thl/test_ledger/test_thl_lm_bp_payout.py | 234 +++-- tests/managers/thl/test_ledger/test_thl_lm_tx.py | 1004 +++++++++++--------- .../test_ledger/test_thl_lm_tx__user_payouts.py | 376 ++++---- tests/managers/thl/test_ledger/test_thl_pem.py | 124 ++- tests/managers/thl/test_ledger/test_user_txs.py | 72 +- tests/managers/thl/test_ledger/test_wallet.py | 36 +- tests/managers/thl/test_maxmind.py | 4 +- tests/managers/thl/test_payout.py | 414 ++++---- tests/managers/thl/test_product.py | 14 +- tests/managers/thl/test_product_prod.py | 12 +- tests/managers/thl/test_session_manager.py | 21 +- tests/managers/thl/test_task_adjustment.py | 6 +- tests/managers/thl/test_task_status.py | 50 +- tests/managers/thl/test_user_manager/test_base.py | 11 +- tests/managers/thl/test_user_manager/test_redis.py | 4 +- .../thl/test_user_manager/test_user_fetch.py | 4 +- .../thl/test_user_manager/test_user_metadata.py | 16 +- tests/managers/thl/test_userhealth.py | 36 +- tests/models/gr/test_authentication.py | 50 +- tests/models/gr/test_business.py | 122 +-- tests/models/gr/test_team.py | 14 +- .../models/legacy/test_user_question_answer_in.py | 16 +- tests/models/test_finance.py | 28 +- tests/models/thl/test_adjustments.py | 2 +- .../thl/test_contest/test_leaderboard_contest.py | 2 +- .../models/thl/test_contest/test_raffle_contest.py | 2 +- tests/models/thl/test_payout.py | 2 +- tests/models/thl/test_product.py | 78 +- tests/models/thl/test_user.py | 6 +- 53 files changed, 2375 insertions(+), 1741 deletions(-) (limited to 'test_utils/models') 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 47ea200eac0eaa7bef02f6ebb05de9afad5ee0d7 Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Wed, 26 Aug 2026 09:17:39 -0700 Subject: Ruff morning! --- generalresearch/models/thl/user.py | 26 +-- test_utils/managers/conftest.py | 100 ++++++----- test_utils/managers/thl/conftest.py | 16 ++ test_utils/models/upk/conftest.py | 11 +- tests/managers/leaderboard.py | 158 ++++++++++------- tests/managers/test_events.py | 65 +++++-- tests/managers/test_lucid.py | 9 +- tests/managers/test_userpid.py | 2 + tests/managers/thl/test_buyer.py | 11 +- tests/managers/thl/test_cashout_method.py | 67 +++++-- tests/managers/thl/test_category.py | 109 +++++++----- tests/managers/thl/test_harmonized_uqa.py | 2 + tests/managers/thl/test_ipinfo.py | 57 ++++-- tests/managers/thl/test_ledger/test_lm_tx_locks.py | 37 ++-- tests/managers/thl/test_ledger/test_wallet.py | 2 +- tests/managers/thl/test_payout.py | 13 +- tests/managers/thl/test_product.py | 45 +++-- tests/managers/thl/test_product_prod.py | 18 +- tests/managers/thl/test_profiling/test_question.py | 25 ++- tests/managers/thl/test_profiling/test_schema.py | 15 +- tests/managers/thl/test_profiling/test_uqa.py | 1 - tests/managers/thl/test_profiling/test_user_upk.py | 20 ++- tests/managers/thl/test_session_manager.py | 51 +++--- tests/managers/thl/test_survey.py | 64 ++++--- tests/managers/thl/test_survey_penalty.py | 15 +- tests/managers/thl/test_task_adjustment.py | 195 +++++++++++---------- tests/managers/thl/test_task_status.py | 111 +++++++----- tests/managers/thl/test_user_manager/test_base.py | 15 +- tests/managers/thl/test_user_manager/test_mysql.py | 27 +-- tests/managers/thl/test_user_manager/test_redis.py | 46 +++-- .../thl/test_user_manager/test_user_fetch.py | 12 +- .../thl/test_user_manager/test_user_metadata.py | 33 +++- tests/managers/thl/test_user_streak.py | 32 +++- tests/managers/thl/test_userhealth.py | 135 +++++++++----- tests/managers/thl/test_wall_manager.py | 58 ++++-- tests/models/custom_types/test_dsn.py | 2 + tests/models/custom_types/test_therest.py | 2 + tests/models/dynata/test_eligbility.py | 2 + tests/models/gr/test_authentication.py | 126 +++++++------ tests/models/gr/test_base.py | 2 + tests/models/gr/test_business.py | 134 +++++++------- tests/models/gr/test_team.py | 14 +- 42 files changed, 1194 insertions(+), 691 deletions(-) delete mode 100644 tests/managers/thl/test_profiling/test_uqa.py (limited to 'test_utils/models') diff --git a/generalresearch/models/thl/user.py b/generalresearch/models/thl/user.py index 11e0d67..d3ffb0d 100644 --- a/generalresearch/models/thl/user.py +++ b/generalresearch/models/thl/user.py @@ -4,7 +4,7 @@ import json import logging import re from datetime import UTC, datetime -from typing import TYPE_CHECKING, Annotated, Self +from typing import TYPE_CHECKING, Annotated, Any, Self from uuid import UUID, uuid4 from pydantic import ( @@ -104,7 +104,7 @@ class User(BaseModel): # --- Validation --- @field_validator("product_user_id") - def check_product_user_id(cls, v: str) -> str: + def check_product_user_id(cls, v: str | None) -> str: if v is not None: if " " in v: raise ValueError("String cannot contain spaces") @@ -122,23 +122,27 @@ class User(BaseModel): # noinspection PyNestedDecoratorsk @field_validator("created", "last_seen") @classmethod - def check_not_in_future(cls, v: AwareDatetime) -> AwareDatetime: + def check_not_in_future(cls, v: AwareDatetime | None) -> AwareDatetime: if v is not None: try: assert v < datetime.now(tz=UTC) - except Exception: + except AssertionError: raise ValueError("Input is in the future") + + assert isinstance(v, AwareDatetime) return v # noinspection PyNestedDecorators @field_validator("created", "last_seen") @classmethod - def check_after_anno_domini(cls, v: AwareDatetime) -> AwareDatetime: + def check_after_anno_domini(cls, v: AwareDatetime | None) -> AwareDatetime: if v is not None: try: assert v > datetime(year=2016, month=7, day=13, tzinfo=UTC) - except Exception: + except AssertionError: raise ValueError("Input is before Anno Domini") + + assert isinstance(v, AwareDatetime) return v @model_validator(mode="after") @@ -168,7 +172,7 @@ class User(BaseModel): ) @classmethod - def is_valid_ubp(cls, *, product_id, product_user_id) -> bool: + def is_valid_ubp(cls, *, product_id: str, product_user_id: str) -> bool: # Attempt to create common_struct solely for validation purposes, # using the product_id and product_user_id try: @@ -178,7 +182,7 @@ class User(BaseModel): product_id=product_id, product_user_id=product_user_id, ) - except Exception as e: + except ValueError as e: logger.info(e) return False else: @@ -186,7 +190,7 @@ class User(BaseModel): # --- Methods --- @staticmethod - def check_bpuid_is_not_bpid(product_id, product_user_id): + def check_bpuid_is_not_bpid(product_id: str | None, product_user_id: str | None): """Unfortunately users were already created failing this constraint, so only check for new users! """ @@ -198,7 +202,7 @@ class User(BaseModel): raise ValueError("product_user_id must not equal the product_id") return True - def to_dict(self) -> dict: + def to_dict(self) -> dict[str, Any]: return self.model_dump(mode="python", exclude={"product"}) def to_json(self) -> str: @@ -291,7 +295,7 @@ class User(BaseModel): # --- Prebuild --- @classmethod - def from_db(cls, res) -> Self: + def from_db(cls, res: dict[str, Any]) -> Self: if res["created"]: res["created"] = res["created"].replace(tzinfo=UTC) if res["last_seen"]: diff --git a/test_utils/managers/conftest.py b/test_utils/managers/conftest.py index 2a9ea00..e5d6015 100644 --- a/test_utils/managers/conftest.py +++ b/test_utils/managers/conftest.py @@ -15,11 +15,17 @@ from generalresearch.managers.gr.team import ( ) from generalresearch.managers.spectrum.survey import SpectrumSurveyManager from generalresearch.managers.thl.buyer import BuyerManager +from generalresearch.managers.thl.cashout_method import ( + CashoutMethodManager, +) from generalresearch.managers.thl.ipinfo import ( GeoIpInfoManager, IPGeonameManager, IPInformationManager, ) +from generalresearch.managers.thl.user_streak import ( + UserStreakManager, +) from generalresearch.managers.thl.userhealth import ( AuditLogManager, IPRecordManager, @@ -114,12 +120,9 @@ def geoipinfo_manager( @pytest.fixture(scope="session") -def cashout_method_manager(thl_web_rw: PostgresConfig): +def cashout_method_manager(thl_web_rw: PostgresConfig) -> CashoutMethodManager: assert thl_web_rw.dsn.path assert "/unittest-" in thl_web_rw.dsn.path - from generalresearch.managers.thl.cashout_method import ( - CashoutMethodManager, - ) return CashoutMethodManager(pg_config=thl_web_rw) @@ -132,12 +135,9 @@ def event_manager(thl_redis_config: RedisConfig): @pytest.fixture(scope="session") -def user_streak_manager(thl_web_rw: PostgresConfig): +def user_streak_manager(thl_web_rw: PostgresConfig) -> UserStreakManager: assert thl_web_rw.dsn.path assert "/unittest-" in thl_web_rw.dsn.path - from generalresearch.managers.thl.user_streak import ( - UserStreakManager, - ) return UserStreakManager(pg_config=thl_web_rw) @@ -169,10 +169,16 @@ def delete_cashoutmethod_db(thl_web_rw: PostgresConfig) -> Callable[..., None]: @pytest.fixture(scope="session") -def setup_cashoutmethod_db(cashout_method_manager, delete_cashoutmethod_db): - delete_cashoutmethod_db() - for x in EXAMPLE_TANGO_CASHOUT_METHODS: - cashout_method_manager.create(x) +def setup_cashoutmethod_db( + cashout_method_manager: CashoutMethodManager, + delete_cashoutmethod_db: Callable[..., None], +) -> Callable[..., None]: + + def _inner(): + delete_cashoutmethod_db() + + for x in EXAMPLE_TANGO_CASHOUT_METHODS: + cashout_method_manager.create(x) # TODO: convert these ids into instances to use. # settings.amt_bonus_cashout_method_id @@ -180,7 +186,9 @@ def setup_cashoutmethod_db(cashout_method_manager, delete_cashoutmethod_db): # cashout_method_manager.create(AMT_ASSIGNMENT_CASHOUT_METHOD) # cashout_method_manager.create(AMT_BONUS_CASHOUT_METHOD) - raise NotImplementedError("Need to implement setup_cashoutmethod_db") + # raise NotImplementedError("Need to implement setup_cashoutmethod_db") + + return _inner # === THL: Marketplaces === @@ -259,33 +267,39 @@ def membership_manager(gr_db: PostgresConfig) -> MembershipManager: @pytest.fixture(scope="session") -def delete_buyers_surveys(thl_web_rw: PostgresConfig, buyer_manager: BuyerManager): - # assert "/unittest-" in thl_web_rw.dsn.path - thl_web_rw.execute_write( - """ - DELETE FROM marketplace_surveystat - WHERE survey_id IN ( - SELECT id - FROM marketplace_survey - WHERE source = %(source)s - );""", - params={"source": Source.TESTING.value}, - ) - thl_web_rw.execute_write( - """ - DELETE FROM marketplace_survey - WHERE buyer_id IN ( - SELECT id - FROM marketplace_buyer - WHERE source = %(source)s - );""", - params={"source": Source.TESTING.value}, - ) - thl_web_rw.execute_write( - """ - DELETE from marketplace_buyer - WHERE source=%(source)s; - """, - params={"source": Source.TESTING.value}, - ) - buyer_manager.populate_caches() +def delete_buyers_surveys( + thl_web_rw: PostgresConfig, buyer_manager: BuyerManager +) -> Callable[..., None]: + + def _inner(): + # assert "/unittest-" in thl_web_rw.dsn.path + thl_web_rw.execute_write( + """ + DELETE FROM marketplace_surveystat + WHERE survey_id IN ( + SELECT id + FROM marketplace_survey + WHERE source = %(source)s + );""", + params={"source": Source.TESTING.value}, + ) + thl_web_rw.execute_write( + """ + DELETE FROM marketplace_survey + WHERE buyer_id IN ( + SELECT id + FROM marketplace_buyer + WHERE source = %(source)s + );""", + params={"source": Source.TESTING.value}, + ) + thl_web_rw.execute_write( + """ + DELETE from marketplace_buyer + WHERE source=%(source)s; + """, + params={"source": Source.TESTING.value}, + ) + buyer_manager.populate_caches() + + return _inner diff --git a/test_utils/managers/thl/conftest.py b/test_utils/managers/thl/conftest.py index 21b2007..d40b7d2 100644 --- a/test_utils/managers/thl/conftest.py +++ b/test_utils/managers/thl/conftest.py @@ -20,6 +20,12 @@ from generalresearch.managers.thl.session import SessionManager from generalresearch.managers.thl.task_adjustment import ( TaskAdjustmentManager, ) +from generalresearch.managers.thl.user_manager.mysql_user_manager import ( + MysqlUserManager, +) +from generalresearch.managers.thl.user_manager.redis_user_manager import ( + RedisUserManager, +) from generalresearch.managers.thl.user_manager.user_manager import ( UserManager, ) @@ -160,6 +166,16 @@ def user_manager( ) +@pytest.fixture(scope="session") +def mysql_user_manager(thl_web_rw: PostgresConfig) -> MysqlUserManager: + return MysqlUserManager(pg_config=thl_web_rw, is_read_replica=False) + + +@pytest.fixture(scope="session") +def redis_user_manager(thl_redis_config: RedisConfig) -> RedisUserManager: + return RedisUserManager(redis_dsn=thl_redis_config) + + @pytest.fixture(scope="session") def user_metadata_manager(thl_web_rw: PostgresConfig) -> UserMetadataManager: assert thl_web_rw.dsn diff --git a/test_utils/models/upk/conftest.py b/test_utils/models/upk/conftest.py index c8855da..ef77dd6 100644 --- a/test_utils/models/upk/conftest.py +++ b/test_utils/models/upk/conftest.py @@ -2,6 +2,7 @@ from __future__ import annotations import os import time +from collections.abc import Callable from typing import TYPE_CHECKING from uuid import UUID @@ -169,9 +170,13 @@ def upk_data( propertymarketplaceassociation_data, propertyitemrange_data, question_data, -) -> None: - # Wait a second to make sure the HarmonizerCache refresh loop pulls these in - time.sleep(2) +) -> Callable[..., None]: + + def _inner(): + # Wait a second to make sure the HarmonizerCache refresh loop pulls these in + time.sleep(2) + + return _inner def test_fixtures(upk_data): diff --git a/tests/managers/leaderboard.py b/tests/managers/leaderboard.py index 3d1818b..d97714d 100644 --- a/tests/managers/leaderboard.py +++ b/tests/managers/leaderboard.py @@ -1,6 +1,9 @@ +from __future__ import annotations + import os import time import zoneinfo +from collections.abc import Callable from datetime import UTC, datetime from decimal import Decimal from uuid import uuid4 @@ -19,10 +22,11 @@ from generalresearch.models.thl.product import ( PayoutConfig, PayoutTransformation, PayoutTransformationPercentArgs, - product: Product, + Product, ) from generalresearch.models.thl.session import Session from generalresearch.models.thl.user import User +from generalresearch.redis_helper import RedisConfig # random uuid for leaderboard tests product_id = uuid4().hex @@ -44,7 +48,9 @@ def session_factory(): def _create_session( - product_user_id="aaa", country_iso="us", user_payout=Decimal("1.00") + product_user_id: str = "aaa", + country_iso: str = "us", + user_payout: Decimal = Decimal("1.00"), ): user = User( product_id=product_id, @@ -74,59 +80,67 @@ def _create_session( @pytest.fixture(scope="function") -def setup_leaderboards(thl_redis): - complete_count = { - "aaa": 10, - "bbb": 6, - "ccc": 6, - "ddd": 6, - "eee": 2, - "fff": 1, - "ggg": 1, - } - sum_payout = {"aaa": 345, "bbb": 100, "ccc": 100} - max_payout = sum_payout - country_iso = "us" - for freq in [ - LeaderboardFrequency.DAILY, - LeaderboardFrequency.WEEKLY, - LeaderboardFrequency.MONTHLY, - ]: - m = LeaderboardManager( - redis_client=thl_redis, - board_code=LeaderboardCode.COMPLETE_COUNT, - freq=freq, - product_id=product_id, - country_iso=country_iso, - within_time=datetime(2025, 2, 5, 12, 12, 12), - ) - thl_redis.delete(m.key) - thl_redis.zadd(m.key, complete_count) - m = LeaderboardManager( - redis_client=thl_redis, - board_code=LeaderboardCode.SUM_PAYOUTS, - freq=freq, - product_id=product_id, - country_iso=country_iso, - within_time=datetime(2025, 2, 5, 12, 12, 12), - ) - thl_redis.delete(m.key) - thl_redis.zadd(m.key, sum_payout) - m = LeaderboardManager( - redis_client=thl_redis, - board_code=LeaderboardCode.LARGEST_PAYOUT, - freq=freq, - product_id=product_id, - country_iso=country_iso, - within_time=datetime(2025, 2, 5, 12, 12, 12), - ) - thl_redis.delete(m.key) - thl_redis.zadd(m.key, max_payout) +def setup_leaderboards(thl_redis: RedisConfig) -> Callable[..., None]: + + def _inner(): + complete_count = { + "aaa": 10, + "bbb": 6, + "ccc": 6, + "ddd": 6, + "eee": 2, + "fff": 1, + "ggg": 1, + } + sum_payout = {"aaa": 345, "bbb": 100, "ccc": 100} + max_payout = sum_payout + country_iso = "us" + for freq in [ + LeaderboardFrequency.DAILY, + LeaderboardFrequency.WEEKLY, + LeaderboardFrequency.MONTHLY, + ]: + m = LeaderboardManager( + redis_client=thl_redis, + board_code=LeaderboardCode.COMPLETE_COUNT, + freq=freq, + product_id=product_id, + country_iso=country_iso, + within_time=datetime(2025, 2, 5, 12, 12, 12, tzinfo=UTC), + ) + thl_redis.delete(m.key) + thl_redis.zadd(m.key, complete_count) + m = LeaderboardManager( + redis_client=thl_redis, + board_code=LeaderboardCode.SUM_PAYOUTS, + freq=freq, + product_id=product_id, + country_iso=country_iso, + within_time=datetime(2025, 2, 5, 12, 12, 12, tzinfo=UTC), + ) + thl_redis.delete(m.key) + thl_redis.zadd(m.key, sum_payout) + m = LeaderboardManager( + redis_client=thl_redis, + board_code=LeaderboardCode.LARGEST_PAYOUT, + freq=freq, + product_id=product_id, + country_iso=country_iso, + within_time=datetime(2025, 2, 5, 12, 12, 12, tzinfo=UTC), + ) + thl_redis.delete(m.key) + thl_redis.zadd(m.key, max_payout) + + return _inner class TestLeaderboards: - def test_leaderboard_manager(self, setup_leaderboards, thl_redis): + def test_leaderboard_manager( + self, setup_leaderboards: Callable[..., None], thl_redis: RedisConfig + ): + setup_leaderboards() + country_iso = "us" board_code = LeaderboardCode.COMPLETE_COUNT freq = LeaderboardFrequency.DAILY @@ -136,7 +150,7 @@ class TestLeaderboards: freq=freq, product_id=product_id, country_iso=country_iso, - within_time=datetime(2025, 2, 5, 0, 0, 0), + within_time=datetime(2025, 2, 5, 0, 0, 0, tzinfo=UTC), ) lb = m.get_leaderboard() assert lb.period_start_local == datetime( @@ -164,7 +178,11 @@ class TestLeaderboards: LeaderboardRow(bpuid="ggg", rank=6, value=1), ] - def test_leaderboard_manager_bpuid(self, setup_leaderboards, thl_redis): + def test_leaderboard_manager_bpuid( + self, setup_leaderboards: Callable[..., None], thl_redis: RedisConfig + ): + setup_leaderboards() + country_iso = "us" board_code = LeaderboardCode.COMPLETE_COUNT freq = LeaderboardFrequency.DAILY @@ -174,7 +192,7 @@ class TestLeaderboards: freq=freq, product_id=product_id, country_iso=country_iso, - within_time=datetime(2025, 2, 5, 12, 12, 12), + within_time=datetime(2025, 2, 5, 12, 12, 12, tzinfo=UTC), ) lb = m.get_leaderboard(bp_user_id="fff", limit=1) @@ -191,7 +209,14 @@ class TestLeaderboards: lb.censor() assert lb.rows[0].bpuid == "ee*" - def test_leaderboard_hit(self, setup_leaderboards, session_factory, thl_redis): + def test_leaderboard_hit( + self, + setup_leaderboards: Callable[..., None], + session_factory: Callable[..., Session], + thl_redis: RedisConfig, + ): + setup_leaderboards() + hit_leaderboards(redis_client=thl_redis, session=session_factory()) for freq in [ @@ -205,7 +230,7 @@ class TestLeaderboards: freq=freq, product_id=product_id, country_iso="us", - within_time=datetime(2025, 2, 5, 12, 12, 12), + within_time=datetime(2025, 2, 5, 12, 12, 12, tzinfo=UTC), ) lb = m.get_leaderboard(limit=1) assert lb.row_count == 7 @@ -216,7 +241,7 @@ class TestLeaderboards: freq=freq, product_id=product_id, country_iso="us", - within_time=datetime(2025, 2, 5, 12, 12, 12), + within_time=datetime(2025, 2, 5, 12, 12, 12, tzinfo=UTC), ) lb = m.get_leaderboard(limit=1) assert lb.row_count == 3 @@ -227,15 +252,20 @@ class TestLeaderboards: freq=freq, product_id=product_id, country_iso="us", - within_time=datetime(2025, 2, 5, 12, 12, 12), + within_time=datetime(2025, 2, 5, 12, 12, 12, tzinfo=UTC), ) lb = m.get_leaderboard(limit=1) assert lb.row_count == 3 assert lb.rows == [LeaderboardRow(bpuid="aaa", rank=1, value=345 + 100)] def test_leaderboard_hit_new_row( - self, setup_leaderboards, session_factory, thl_redis + self, + setup_leaderboards: Callable[..., None], + session_factory: Callable[..., None], + thl_redis: RedisConfig, ): + setup_leaderboards() + session = session_factory(product_user_id="zzz") hit_leaderboards(redis_client=thl_redis, session=session) m = LeaderboardManager( @@ -244,24 +274,20 @@ class TestLeaderboards: freq=LeaderboardFrequency.DAILY, product_id=product_id, country_iso="us", - within_time=datetime(2025, 2, 5, 12, 12, 12), + within_time=datetime(2025, 2, 5, 12, 12, 12, tzinfo=UTC), ) lb = m.get_leaderboard() assert lb.row_count == 8 assert LeaderboardRow(bpuid="zzz", value=1, rank=6) in lb.rows - def test_leaderboard_country(self, thl_redis): + def test_leaderboard_country(self, thl_redis: RedisConfig): m = LeaderboardManager( redis_client=thl_redis, board_code=LeaderboardCode.COMPLETE_COUNT, freq=LeaderboardFrequency.DAILY, product_id=product_id, country_iso="jp", - within_time=datetime( - 2025, - 2, - 1, - ), + within_time=datetime(2025, 2, 1, tzinfo=UTC), ) lb = m.get_leaderboard() assert lb.row_count == 0 diff --git a/tests/managers/test_events.py b/tests/managers/test_events.py index a6d3a6b..cb32275 100644 --- a/tests/managers/test_events.py +++ b/tests/managers/test_events.py @@ -1,6 +1,9 @@ +from __future__ import annotations + import math import random import time +from collections.abc import Callable from datetime import UTC, datetime, timedelta from decimal import Decimal from functools import partial @@ -9,7 +12,8 @@ from uuid import uuid4 import pytest -from generalresearch.managers.events import EventSubscriber +from generalresearch.managers.events import EventManager, EventSubscriber +from generalresearch.managers.thl.product import ProductManager from generalresearch.models import Source from generalresearch.models.events import ( AggregateBySource, @@ -21,22 +25,23 @@ from generalresearch.models.legacy.bucket import Bucket from generalresearch.models.thl.definitions import Status, StatusCode1 from generalresearch.models.thl.session import Session, Wall from generalresearch.models.thl.user import User +from generalresearch.redis_helper import RedisConfig # We don't need anything in the db, so not using the db fixtures @pytest.fixture(scope="function") -def product_id(product_manager): +def product_id(product_manager: ProductManager) -> str: return uuid4().hex @pytest.fixture(scope="function") -def user_factory(product_id): +def user_factory(product_id: str): return partial(create_dummy, product_id=product_id) @pytest.fixture(scope="function") -def event_subscriber(thl_redis_config: RedisConfig, product_id): - return EventSubscriber(redis_config=thl_redis_config: RedisConfig, product_id=product_id) +def event_subscriber(thl_redis_config: RedisConfig, product_id: str) -> EventSubscriber: + return EventSubscriber(redis_config=thl_redis_config, product_id=product_id) def create_dummy( @@ -53,7 +58,7 @@ def create_dummy( class TestActiveUsers: - def test_run_empty(self, event_manager, product_id): + def test_run_empty(self, event_manager: EventManager, product_id: str): res = event_manager.get_user_stats(product_id) assert res == { "active_users_last_1h": 0, @@ -62,7 +67,12 @@ class TestActiveUsers: "in_progress_users": 0, } - def test_run(self, event_manager, product_id, user_factory): + def test_run( + self, + event_manager: EventManager, + product_id: str, + user_factory: Callable[..., User], + ): event_manager.clear_global_user_stats() user1: User = user_factory() @@ -89,6 +99,8 @@ class TestActiveUsers: # Create a 2nd user in another product product_id2 = uuid4().hex user2: User = user_factory(product_id=product_id2) + assert isinstance(user2, User) + assert isinstance(user2.created, datetime) # Change to say user was created >24 hrs ago user2.created = user2.created - timedelta(hours=25) event_manager.handle_user(user2) @@ -115,7 +127,12 @@ class TestActiveUsers: "in_progress_users": 0, } - def test_inprogress(self, event_manager, product_id, user_factory): + def test_inprogress( + self, + event_manager: EventSubscriber, + product_id: str, + user_factory: Callable[..., User], + ): event_manager.clear_global_user_stats() user1: User = user_factory() user2: User = user_factory() @@ -138,7 +155,12 @@ class TestActiveUsers: res = event_manager.get_user_stats(product_id) assert res["in_progress_users"] == 1 - def test_expiry(self, event_manager, product_id, user_factory): + def test_expiry( + self, + event_manager: EventManager, + product_id: str, + user_factory: Callable[..., User], + ): event_manager.clear_global_user_stats() user1: User = user_factory() event_manager.handle_user(user1) @@ -166,7 +188,7 @@ class TestActiveUsers: class TestSessionStats: - def test_run_empty(self, event_manager, product_id): + def test_run_empty(self, event_manager: EventManager, product_id: str): res = event_manager.get_session_stats(product_id) assert res == { "session_enters_last_1h": 0, @@ -185,7 +207,14 @@ class TestSessionStats: "session_fail_avg_loi_last_24h": None, } - def test_run(self, event_manager, product_id, user_factory: Callable[..., User], utc_now, utc_hour_ago): + def test_run( + self, + event_manager: EventManager, + product_id: str, + user_factory: Callable[..., User], + utc_now: datetime, + utc_hour_ago: datetime, + ): event_manager.clear_global_session_stats() user: User = user_factory() @@ -306,7 +335,7 @@ class TestSessionStats: class TestTaskStatsManager: - def test_empty(self, event_manager): + def test_empty(self, event_manager: EventManager): event_manager.clear_task_stats() assert event_manager.get_task_stats_raw() == { "live_task_count": AggregateBySource(total=0), @@ -320,7 +349,7 @@ class TestTaskStatsManager: assert sm.data.task_created_count_last_24h.total == 0 assert sm.data.live_tasks_max_payout.value is None - def test(self, event_manager): + def test(self, event_manager: EventManager): event_manager.clear_task_stats() event_manager.set_source_task_stats( source=Source.TESTING, @@ -445,12 +474,12 @@ class TestTaskStatsManager: class TestChannelsSubscriptions: def test_stats_worker( self, - event_manager, - event_subscriber, - product_id, + event_manager: EventManager, + event_subscriber: EventSubscriber, + product_id: str, user_factory: Callable[..., User], - utc_hour_ago, - utc_now, + utc_hour_ago: datetime, + utc_now: datetime, ): event_manager.clear_stats() assert event_subscriber.pubsub diff --git a/tests/managers/test_lucid.py b/tests/managers/test_lucid.py index 654b58d..20dca22 100644 --- a/tests/managers/test_lucid.py +++ b/tests/managers/test_lucid.py @@ -1,6 +1,9 @@ +from __future__ import annotations + import pytest from generalresearch.managers.lucid.profiling import get_profiling_library +from generalresearch.pg_helper import PostgresConfig qids = ["42", "43", "45", "97", "120", "639", "15297"] @@ -8,9 +11,9 @@ qids = ["42", "43", "45", "97", "120", "639", "15297"] class TestLucidProfiling: @pytest.mark.skip - def test_get_library(self, thl_web_rr): + def test_get_library(self, thl_web_rr: PostgresConfig): pks = [(qid, "us", "eng") for qid in qids] - qs = get_profiling_library(thl_web_rr: PostgresConfig, pks=pks) + qs = get_profiling_library(thl_web_rr, pks=pks) assert len(qids) == len(qs) # just making sure this doesn't raise errors @@ -19,5 +22,5 @@ class TestLucidProfiling: # a lot will fail parsing because they have no options or the options are blank # just asserting that we get some back - qs = get_profiling_library(thl_web_rr: PostgresConfig, country_iso="mx", language_iso="spa") + qs = get_profiling_library(thl_web_rr, country_iso="mx", language_iso="spa") assert len(qs) > 100 diff --git a/tests/managers/test_userpid.py b/tests/managers/test_userpid.py index 36c2de9..e74e40b 100644 --- a/tests/managers/test_userpid.py +++ b/tests/managers/test_userpid.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import pytest from pydantic import MySQLDsn diff --git a/tests/managers/thl/test_buyer.py b/tests/managers/thl/test_buyer.py index 69ea105..6776ab3 100644 --- a/tests/managers/thl/test_buyer.py +++ b/tests/managers/thl/test_buyer.py @@ -1,3 +1,8 @@ +from __future__ import annotations + +from collections.abc import Callable + +from generalresearch.managers.thl.buyer import BuyerManager from generalresearch.models import Source @@ -5,10 +10,12 @@ class TestBuyer: def test( self, - delete_buyers_surveys, - buyer_manager, + delete_buyers_surveys: Callable[..., None], + buyer_manager: BuyerManager, ): + delete_buyers_surveys() + bs = buyer_manager.bulk_get_or_create(source=Source.TESTING, codes=["a", "b"]) assert len(bs) == 2 buyer_a = bs[0] diff --git a/tests/managers/thl/test_cashout_method.py b/tests/managers/thl/test_cashout_method.py index ee52188..451d3e0 100644 --- a/tests/managers/thl/test_cashout_method.py +++ b/tests/managers/thl/test_cashout_method.py @@ -1,5 +1,14 @@ +from __future__ import annotations + +from collections.abc import Callable + import pytest +from generalresearch.config import GRLBaseSettings +from generalresearch.managers.thl.cashout_method import ( + CashoutMethodManager, +) +from generalresearch.models.thl.user import User from generalresearch.models.thl.wallet import PayoutType from generalresearch.models.thl.wallet.cashout_method import ( CashMailCashoutMethodData, @@ -13,15 +22,26 @@ from test_utils.managers.cashout_methods import ( class TestTangoCashoutMethods: - def test_create_and_get(self, cashout_method_manager, setup_cashoutmethod_db): + def test_create_and_get( + self, + cashout_method_manager: CashoutMethodManager, + setup_cashoutmethod_db: Callable[..., None], + ): + setup_cashoutmethod_db() + res = cashout_method_manager.filter(payout_types=[PayoutType.TANGO]) assert len(res) == 2 - cm = [x for x in res if x.ext_id == "U025035"][0] + cm = next(x for x in res if x.ext_id == "U025035") assert EXAMPLE_TANGO_CASHOUT_METHODS[0] == cm def test_user( - self, cashout_method_manager, user_with_wallet, setup_cashoutmethod_db + self, + cashout_method_manager: CashoutMethodManager, + user_with_wallet: User, + setup_cashoutmethod_db: Callable[..., None], ): + setup_cashoutmethod_db() + res = cashout_method_manager.get_cashout_methods(user_with_wallet) # This user ONLY has the two tango cashout methods, no AMT assert len(res) == 2 @@ -29,19 +49,31 @@ class TestTangoCashoutMethods: class TestAMTCashoutMethods: - def test_create_and_get(self, cashout_method_manager, setup_cashoutmethod_db): + def test_create_and_get( + self, + settings: GRLBaseSettings, + cashout_method_manager: CashoutMethodManager, + setup_cashoutmethod_db: Callable[..., None], + ): + setup_cashoutmethod_db() + res = cashout_method_manager.filter(payout_types=[PayoutType.AMT]) assert len(res) == 2 - cm = [x for x in res if x.name == "AMT Assignment"][0] - assert AMT_ASSIGNMENT_CASHOUT_METHOD == cm + cm = next(x for x in res if x.name == "AMT Assignment") + assert settings.amt_assignment_cashout_method_id == cm - cm = [x for x in res if x.name == "AMT Bonus"][0] - assert AMT_BONUS_CASHOUT_METHOD == cm + cm = next(x for x in res if x.name == "AMT Bonus") + assert settings.amt_bonus_cashout_method_id == cm def test_user( - self, cashout_method_manager, user_with_wallet_amt, setup_cashoutmethod_db + self, + cashout_method_manager: CashoutMethodManager, + user_with_wallet_amt: User, + setup_cashoutmethod_db: Callable[..., None], ): + setup_cashoutmethod_db() + res = cashout_method_manager.get_cashout_methods(user_with_wallet_amt) # This user has the 2 tango, plus amt bonus & assignment assert len(res) == 4 @@ -49,14 +81,22 @@ class TestAMTCashoutMethods: class TestUserCashoutMethods: - def test(self, cashout_method_manager, user_with_wallet, delete_cashoutmethod_db): + def test( + self, + cashout_method_manager: CashoutMethodManager, + user_with_wallet: User, + delete_cashoutmethod_db: Callable[..., None], + ): delete_cashoutmethod_db() res = cashout_method_manager.get_cashout_methods(user_with_wallet) assert len(res) == 0 def test_cash_in_mail( - self, cashout_method_manager, user_with_wallet, delete_cashoutmethod_db + self, + cashout_method_manager: CashoutMethodManager, + user_with_wallet: User, + delete_cashoutmethod_db: Callable[..., None], ): delete_cashoutmethod_db() @@ -95,7 +135,10 @@ class TestUserCashoutMethods: assert len(res) == 2 def test_paypal( - self, cashout_method_manager, user_with_wallet, delete_cashoutmethod_db + self, + cashout_method_manager: CashoutMethodManager, + user_with_wallet: User, + delete_cashoutmethod_db: Callable[..., None], ): delete_cashoutmethod_db() diff --git a/tests/managers/thl/test_category.py b/tests/managers/thl/test_category.py index ad0f07b..ec52aae 100644 --- a/tests/managers/thl/test_category.py +++ b/tests/managers/thl/test_category.py @@ -1,12 +1,18 @@ +from __future__ import annotations + +from collections.abc import Callable + import pytest +from generalresearch.managers.thl.category import CategoryManager from generalresearch.models.thl.category import Category +from generalresearch.pg_helper import PostgresConfig class TestCategory: @pytest.fixture - def beauty_fitness(self, thl_web_rw): + def beauty_fitness(self, thl_web_rw: PostgresConfig) -> Category: return Category( uuid="12c1e96be82c4642a07a12a90ce6f59e", @@ -16,72 +22,83 @@ class TestCategory: ) @pytest.fixture - def hair_care(self, beauty_fitness): + def hair_care(self, beauty_fitness: Category) -> Category: return Category( uuid="dd76c4b565d34f198dad3687326503d6", adwords_vertical_id="146", label="Hair Care", - path="/Beauty & Fitness/Hair Care", + path=f"{beauty_fitness.path}/Hair Care", ) @pytest.fixture - def hair_loss(self, hair_care): + def hair_loss(self, hair_care: Category) -> Category: return Category( uuid="aacff523c8e246888215611ec3b823c0", adwords_vertical_id="235", label="Hair Loss", - path="/Beauty & Fitness/Hair Care/Hair Loss", + path=f"{hair_care}/Hair Loss", ) @pytest.fixture def category_data( - self, category_manager, thl_web_rw, beauty_fitness, hair_care, hair_loss - ): - cats = [beauty_fitness, hair_care, hair_loss] - data = [x.model_dump(mode="json") for x in cats] - # We need the parent pk's to set the parent_id. So insert all without a parent, - # then pull back all pks and map to the parents as parsed by the parent_path - query = """ - INSERT INTO marketplace_category - (uuid, adwords_vertical_id, label, path) - VALUES - (%(uuid)s, %(adwords_vertical_id)s, %(label)s, %(path)s) - ON CONFLICT (uuid) DO NOTHING; - """ - with thl_web_rw.make_connection() as conn: - with conn.cursor() as c: - c.executemany(query=query, params_seq=data) - conn.commit() - - res = thl_web_rw.execute_sql_query("SELECT id, path FROM marketplace_category") - path_id = {x["path"]: x["id"] for x in res} - data = [ - {"id": path_id[c.path], "parent_id": path_id[c.parent_path]} - for c in cats - if c.parent_path - ] - query = """ - UPDATE marketplace_category - SET parent_id = %(parent_id)s - WHERE id = %(id)s; - """ - with thl_web_rw.make_connection() as conn: - with conn.cursor() as c: - c.executemany(query=query, params_seq=data) - conn.commit() - - category_manager.populate_caches() + self, + category_manager: CategoryManager, + thl_web_rw: PostgresConfig, + beauty_fitness: Category, + hair_care: Category, + hair_loss: Category, + ) -> Callable[..., None]: + + def _inner(): + cats = [beauty_fitness, hair_care, hair_loss] + data = [x.model_dump(mode="json") for x in cats] + # We need the parent pk's to set the parent_id. So insert all without a parent, + # then pull back all pks and map to the parents as parsed by the parent_path + query = """ + INSERT INTO marketplace_category + (uuid, adwords_vertical_id, label, path) + VALUES + (%(uuid)s, %(adwords_vertical_id)s, %(label)s, %(path)s) + ON CONFLICT (uuid) DO NOTHING; + """ + with thl_web_rw.make_connection() as conn: + with conn.cursor() as c: + c.executemany(query=query, params_seq=data) + conn.commit() + + res = thl_web_rw.execute_sql_query( + "SELECT id, path FROM marketplace_category" + ) + path_id = {x["path"]: x["id"] for x in res} + data = [ + {"id": path_id[c.path], "parent_id": path_id[c.parent_path]} + for c in cats + if c.parent_path + ] + query = """ + UPDATE marketplace_category + SET parent_id = %(parent_id)s + WHERE id = %(id)s; + """ + with thl_web_rw.make_connection() as conn: + with conn.cursor() as c: + c.executemany(query=query, params_seq=data) + conn.commit() + + category_manager.populate_caches() + + return _inner def test( self, - category_data, - category_manager, - beauty_fitness, - hair_care, - hair_loss, + category_data: Callable[..., None], + category_manager: CategoryManager, + beauty_fitness: Category, ): + category_data() + # category_manager on init caches the category info. This rarely/never changes so this is fine, # but now that tests get run on a new db each time, the category_manager is inited before # the fixtures run. so category_manager's cache needs to be rerun diff --git a/tests/managers/thl/test_harmonized_uqa.py b/tests/managers/thl/test_harmonized_uqa.py index 81ac080..84eeb56 100644 --- a/tests/managers/thl/test_harmonized_uqa.py +++ b/tests/managers/thl/test_harmonized_uqa.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from datetime import UTC, datetime import pytest diff --git a/tests/managers/thl/test_ipinfo.py b/tests/managers/thl/test_ipinfo.py index c89312b..48b9efd 100644 --- a/tests/managers/thl/test_ipinfo.py +++ b/tests/managers/thl/test_ipinfo.py @@ -1,3 +1,5 @@ +from collections.abc import Callable + import faker from generalresearch.managers.thl.ipinfo import ( @@ -5,14 +7,22 @@ from generalresearch.managers.thl.ipinfo import ( IPGeonameManager, IPInformationManager, ) -from generalresearch.models.thl.ipinfo import IPGeoname, IPInformation +from generalresearch.models.thl.ipinfo import ( + GeoIPInformation, + IPGeoname, + IPInformation, +) +from generalresearch.pg_helper import PostgresConfig +from generalresearch.redis_helper import RedisConfig fake = faker.Faker() class TestIPGeonameManager: - def test_init(self, thl_web_rr: PostgresConfig, ip_geoname_manager: IPGeonameManager): + def test_init( + self, thl_web_rr: PostgresConfig, ip_geoname_manager: IPGeonameManager + ): instance = IPGeonameManager(pg_config=thl_web_rr) assert isinstance(instance, IPGeonameManager) @@ -31,7 +41,9 @@ class TestIPGeonameManager: class TestIPInformationManager: - def test_init(self, thl_web_rr: PostgresConfig, ip_information_manager: IPInformationManager): + def test_init( + self, thl_web_rr: PostgresConfig, ip_information_manager: IPInformationManager + ): instance = IPInformationManager(pg_config=thl_web_rr) assert isinstance(instance, IPInformationManager) assert isinstance(ip_information_manager, IPInformationManager) @@ -45,7 +57,12 @@ class TestIPInformationManager: assert res[0].model_dump_json() == instance.model_dump_json() - def test_prefetch_geoname(self, ip_information, ip_geoname, thl_web_rr): + def test_prefetch_geoname( + self, + ip_information: IPInformation, + ip_geoname: IPGeoname, + thl_web_rr: PostgresConfig, + ): assert isinstance(ip_information, IPInformation) assert ip_information.geoname_id == ip_geoname.geoname_id @@ -62,11 +79,16 @@ class TestGeoIpInfoManager: thl_redis_config: RedisConfig, geoipinfo_manager: GeoIpInfoManager, ): - instance = GeoIpInfoManager(pg_config=thl_web_rr: PostgresConfig, redis_config=thl_redis_config) + instance = GeoIpInfoManager(pg_config=thl_web_rr, redis_config=thl_redis_config) assert isinstance(instance, GeoIpInfoManager) assert isinstance(geoipinfo_manager, GeoIpInfoManager) - def test_multi(self, ip_information_factory, ip_geoname, geoipinfo_manager): + def test_multi( + self, + ip_information_factory: Callable[..., IPInformation], + ip_geoname: IPGeoname, + geoipinfo_manager: GeoIpInfoManager, + ): ip = fake.ipv4_public() ip_information_factory(ip=ip, geoname=ip_geoname) ips = [ip] @@ -93,7 +115,12 @@ class TestGeoIpInfoManager: assert res[ip] is not None assert res[ip2] is not None - def test_multi_ipv6(self, ip_information_factory, ip_geoname, geoipinfo_manager): + def test_multi_ipv6( + self, + ip_information_factory: Callable[..., IPInformation], + ip_geoname: IPGeoname, + geoipinfo_manager: GeoIpInfoManager, + ): ip = fake.ipv6() # Make another IP that will be in the same /64 block. ip2 = ip[:-1] + "a" if ip[-1] != "a" else ip[:-1] + "b" @@ -108,13 +135,19 @@ class TestGeoIpInfoManager: # Looks up in redis, if not exists, looks in mysql, then sets # the caches that didn't exist. res = geoipinfo_manager.get_multi(ip_addresses=ips) - assert res[ip].ip == ip - assert res[ip].lookup_prefix == "/64" - assert res[ip2].ip == ip2 - assert res[ip2].lookup_prefix == "/64" + + res1 = res[ip] + assert isinstance(res1, GeoIPInformation) + assert res1.ip == ip + assert res1.lookup_prefix == "/64" + + res2 = res[ip2] + assert isinstance(res2, GeoIPInformation) + assert res2.ip == ip2 + assert res2.lookup_prefix == "/64" # they should be the same basically, except for the ip - def test_doesnt_exist(self, geoipinfo_manager): + def test_doesnt_exist(self, geoipinfo_manager: GeoIpInfoManager): ip = fake.ipv4_public() res = geoipinfo_manager.get_multi(ip_addresses=[ip]) assert res == {ip: None} diff --git a/tests/managers/thl/test_ledger/test_lm_tx_locks.py b/tests/managers/thl/test_ledger/test_lm_tx_locks.py index 9158e15..e603632 100644 --- a/tests/managers/thl/test_ledger/test_lm_tx_locks.py +++ b/tests/managers/thl/test_ledger/test_lm_tx_locks.py @@ -1,9 +1,10 @@ from __future__ import annotations import logging -from collections.abc import Callable +from collections.abc import Callable, Generator from datetime import UTC, datetime, timedelta from decimal import Decimal +from logging import LogCaptureFixture import pytest @@ -41,7 +42,7 @@ class TestLedgerLocks: session_factory: Callable[..., Session], product_user_wallet_no: Product, create_main_accounts: Callable[..., None], - caplog, + caplog: Generator[LogCaptureFixture], thl_ledger_manager: ThlLedgerManager, ledger_manager: LedgerManager, utc_hour_ago: datetime, @@ -126,14 +127,15 @@ class TestLedgerLocks: # purposely hold the lock open tx = None ledger_manager.redis_client.set(lock_name, "1") - with caplog.at_level(logging.ERROR): - with pytest.raises(expected_exception=LedgerTransactionCreateLockError): - tx = thl_ledger_manager.create_tx_protected( - lock_key=lock_key, - condition=condition, - create_tx_func=create_tx_func, - ) - assert tx is None + with caplog.at_level(logging.ERROR), pytest.raises( + expected_exception=LedgerTransactionCreateLockError + ): + tx = thl_ledger_manager.create_tx_protected( + lock_key=lock_key, + condition=condition, + create_tx_func=create_tx_func, + ) + assert tx is None assert "Unable to acquire lock within the time specified" in caplog.text ledger_manager.redis_client.delete(lock_name) @@ -143,7 +145,7 @@ class TestLedgerLocks: product_user_wallet_no: Product, create_main_accounts: Callable[..., None], delete_ledger_db: Callable[..., None], - caplog, + caplog: Generator[LogCaptureFixture], thl_ledger_manager: ThlLedgerManager, ledger_manager: LedgerManager, ): @@ -226,12 +228,13 @@ class TestLedgerLocks: # Purposely hold the lock open ledger_manager.redis_client.set(name=lock_name, value="1") - with caplog.at_level(logging.DEBUG): - with pytest.raises(expected_exception=LedgerTransactionCreateLockError): - tx = thl_ledger_manager.create_tx_task_complete( - wall=wall3, user=user, created=wall3.started - ) - assert isinstance(tx, LedgerTransaction) + with caplog.at_level(logging.DEBUG), pytest.raises( + expected_exception=LedgerTransactionCreateLockError + ): + tx = thl_ledger_manager.create_tx_task_complete( + wall=wall3, user=user, created=wall3.started + ) + assert isinstance(tx, LedgerTransaction) assert "Unable to acquire lock within the time specified" in caplog.text # Release the lock diff --git a/tests/managers/thl/test_ledger/test_wallet.py b/tests/managers/thl/test_ledger/test_wallet.py index 9e886db..cad3ea4 100644 --- a/tests/managers/thl/test_ledger/test_wallet.py +++ b/tests/managers/thl/test_ledger/test_wallet.py @@ -6,7 +6,6 @@ from uuid import uuid4 import pytest -from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager from generalresearch.managers.thl.product import ProductManager from generalresearch.models.thl.product import ( @@ -55,6 +54,7 @@ class TestGetUserWalletBalance: user: User = user_factory(schrute_product) balance = thl_ledger_manager.get_user_wallet_balance(user=user) assert balance == 0 + assert isinstance(user.product, Product) balance_string = user.product.format_payout_format(Decimal(balance) / 100) assert balance_string == "0 Schrute Bucks" redeemable_balance = thl_ledger_manager.get_user_redeemable_wallet_balance( diff --git a/tests/managers/thl/test_payout.py b/tests/managers/thl/test_payout.py index 153bee9..e6c597b 100644 --- a/tests/managers/thl/test_payout.py +++ b/tests/managers/thl/test_payout.py @@ -513,7 +513,7 @@ class TestBusinessPayoutEventManager: assert int(res.deduction.sum()) == available_balance - 1 # Slightly less - with pytest.raises(expected_exception=Exception) as cm: + with pytest.raises(expected_exception=ValueError): res = business_payout_event_manager.recoup_proportional( df=df, target_amount=available_balance + 1 ) @@ -731,12 +731,9 @@ class TestBusinessPayoutEventManager: def test_ach_payment( self, - product: Product, mnt_filepath: GRLDatasets, thl_ledger_manager: ThlLedgerManager, client_no_amm: DaskClient, - thl_redis_config: RedisConfig, - payout_event_manager: PayoutEventManager, brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager, business_payout_event_manager: BusinessPayoutEventManager, delete_ledger_db: Callable[..., None], @@ -899,8 +896,6 @@ class TestBusinessPayoutEventManager: pop_ledger=pop_ledger_merge, ) business.prebuild_payouts( - thl_pg_config=thl_web_rr, - thl_lm=thl_ledger_manager, bpem=business_payout_event_manager, ) assert isinstance(business.payouts, list) @@ -930,13 +925,10 @@ class TestBusinessPayoutEventManager: def test_ach_payment_partial_amount( self, - product: Product, mnt_filepath: GRLDatasets, thl_ledger_manager: ThlLedgerManager, client_no_amm: DaskClient, - thl_redis_config: RedisConfig, payout_event_manager: PayoutEventManager, - brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager, business_payout_event_manager: BusinessPayoutEventManager, delete_ledger_db: Callable[..., None], create_main_accounts: Callable[..., None], @@ -948,8 +940,6 @@ class TestBusinessPayoutEventManager: session_with_tx_factory: Callable[..., None], pop_ledger_merge: PopLedgerMerge, start: datetime, - bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], - adj_to_fail_with_tx_factory: Callable[..., None], thl_web_rr: PostgresConfig, ledger_manager: LedgerManager, product_manager: ProductManager, @@ -1076,7 +1066,6 @@ class TestBusinessPayoutEventManager: thl_ledger_manager: ThlLedgerManager, client_no_amm: DaskClient, payout_event_manager: PayoutEventManager, - brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager, business_payout_event_manager: BusinessPayoutEventManager, delete_ledger_db: Callable[..., None], create_main_accounts: Callable[..., None], diff --git a/tests/managers/thl/test_product.py b/tests/managers/thl/test_product.py index 31b7b73..f93ac36 100644 --- a/tests/managers/thl/test_product.py +++ b/tests/managers/thl/test_product.py @@ -1,10 +1,15 @@ +from __future__ import annotations + +from collections.abc import Callable from uuid import uuid4 import pytest +from generalresearch.managers.thl.product import ProductManager from generalresearch.models import Source +from generalresearch.models.gr.team import Team from generalresearch.models.thl.product import ( - product: Product, + Product, ProfilingConfig, SourceConfig, SourcesConfig, @@ -16,7 +21,7 @@ from generalresearch.models.thl.product import ( class TestProductManagerGetMethods: - def test_get_by_uuid(self, product_manager): + def test_get_by_uuid(self, product_manager: ProductManager): product: Product = product_manager.create_dummy( product_id=uuid4().hex, team_id=uuid4().hex, @@ -36,10 +41,10 @@ class TestProductManagerGetMethods: product_manager.get_by_uuid(product_uuid=uuid4().hex) assert "product not found" in str(cm.value) - def test_get_by_uuids(self, product_manager): + def test_get_by_uuids(self, product_manager: ProductManager): cnt = 5 - product_uuids = [uuid4().hex for idx in range(cnt)] + product_uuids = [uuid4().hex for _ in range(cnt)] for product_id in product_uuids: product_manager.create_dummy( product_id=product_id, @@ -61,7 +66,7 @@ class TestProductManagerGetMethods: product_manager.get_by_uuids(product_uuids=product_uuids + ["abc123"]) assert "invalid uuid" in str(cm.value) - def test_get_by_uuid_if_exists(self, product_manager): + def test_get_by_uuid_if_exists(self, product_manager: ProductManager): product: Product = product_manager.create_dummy( product_id=uuid4().hex, team_id=uuid4().hex, @@ -73,7 +78,7 @@ class TestProductManagerGetMethods: instance = product_manager.get_by_uuid_if_exists(product_uuid="abc123") assert instance == None - def test_get_by_uuids_if_exists(self, product_manager): + def test_get_by_uuids_if_exists(self, product_manager: ProductManager): product_uuids = [uuid4().hex for _ in range(2)] for product_id in product_uuids: product_manager.create_dummy( @@ -105,8 +110,8 @@ class TestProductManagerGetMethods: # for instance in res: # assert isinstance(instance, Product) - def test_get_by_business_ids(self, product_manager): - business_ids = [uuid4().hex for i in range(5)] + def test_get_by_business_ids(self, product_manager: ProductManager): + business_ids = [uuid4().hex for _ in range(5)] product_manager.fetch_uuids(business_uuids=business_ids) @@ -123,7 +128,7 @@ class TestProductManagerGetMethods: class TestProductManagerCreation: - def test_base(self, product_manager): + def test_base(self, product_manager: ProductManager): instance = product_manager.create_dummy( product_id=uuid4().hex, team_id=uuid4().hex, @@ -135,7 +140,7 @@ class TestProductManagerCreation: class TestProductManagerCreate: - def test_create_simple(self, product_manager): + def test_create_simple(self, product_manager: ProductManager): # Always required: product_id, team_id, name, redirect_url # Required internally - if not passed use default: harmonizer_domain, # commission_pct, sources @@ -181,9 +186,9 @@ class TestProductManager: def test_get_by_uuid1( self, product_manager: ProductManager, - team, + team: Team, product: Product, - product_factory, + product_factory: Callable[..., Product], ): p1 = product_factory(team=team) instance = product_manager.get_by_uuid(product_uuid=p1.uuid) @@ -209,7 +214,9 @@ class TestProductManager: assert 0 == instance.user_create_config.min_hourly_create_limit assert instance.user_create_config.max_hourly_create_limit is None - def test_get_by_uuid3(self, product_manager: ProductManager, product_factory): + def test_get_by_uuid3( + self, product_manager: ProductManager, product_factory: Callable[..., Product] + ): p3 = product_factory() instance = product_manager.get_by_uuid(p3.id) assert instance.id == p3.id @@ -225,7 +232,7 @@ class TestProductManager: assert instance.user_create_config.max_hourly_create_limit is None assert not instance.user_wallet_config.enabled - def test_sources(self, product_manager): + def test_sources(self, product_manager: ProductManager): user_defined = [SourceConfig(name=Source.DYNATA, active=False)] sources_config = SourcesConfig(user_defined=user_defined) p = product_manager.create_dummy(sources_config=sources_config) @@ -240,7 +247,7 @@ class TestProductManager: assert not dynata.active assert all(x.active is True for x in p2.sources if x.name != Source.DYNATA) - def test_global_sources(self, product_manager): + def test_global_sources(self, product_manager: ProductManager): sources_config = SupplyConfig( policies=[ SupplyPolicy( @@ -267,7 +274,7 @@ class TestProductManager: p2 = product_manager.get_by_uuid(p1.id) assert p1 == p2 - def test_user_health_config(self, product_manager): + def test_user_health_config(self, product_manager: ProductManager): p = product_manager.create_dummy( user_health_config=UserHealthConfig(banned_countries=["ng", "in"]) ) @@ -278,7 +285,7 @@ class TestProductManager: assert p2.user_health_config.banned_countries == ["in", "ng"] assert p2.user_health_config.allow_ban_iphist - def test_profiling_config(self, product_manager): + def test_profiling_config(self, product_manager: ProductManager): p = product_manager.create_dummy( profiling_config=ProfilingConfig(max_questions=1) ) @@ -325,7 +332,7 @@ class TestProductManager: class TestProductManagerUpdate: - def test_update(self, product_manager): + def test_update(self, product_manager: ProductManager): p = product_manager.create_dummy() p.name = "new name" p.enabled = False @@ -346,7 +353,7 @@ class TestProductManagerUpdate: class TestProductManagerCacheClear: - def test_cache_clear(self, product_manager): + def test_cache_clear(self, product_manager: ProductManager): p = product_manager.create_dummy() product_manager.get_by_uuid(product_uuid=p.id) product_manager.get_by_uuid(product_uuid=p.id) diff --git a/tests/managers/thl/test_product_prod.py b/tests/managers/thl/test_product_prod.py index 0f622b6..8734210 100644 --- a/tests/managers/thl/test_product_prod.py +++ b/tests/managers/thl/test_product_prod.py @@ -1,8 +1,12 @@ +from __future__ import annotations + import logging +from collections.abc import Callable from uuid import uuid4 import pytest +from generalresearch.managers.thl.product import ProductManager from generalresearch.models.thl.product import Product logger = logging.getLogger() @@ -10,7 +14,9 @@ logger = logging.getLogger() class TestProductManagerGetMethods: - def test_get_by_uuid(self, product_manager: ProductManager, product_factory): + def test_get_by_uuid( + self, product_manager: ProductManager, product_factory: Callable[..., Product] + ): # Just test that we load properly for p in [product_factory(), product_factory(), product_factory()]: instance = product_manager.get_by_uuid(product_uuid=p.id) @@ -22,7 +28,9 @@ class TestProductManagerGetMethods: product_manager.get_by_uuid(product_uuid=uuid4().hex) assert "product not found" in str(cm.value) - def test_get_by_uuids(self, product_manager: ProductManager, product_factory): + def test_get_by_uuids( + self, product_manager: ProductManager, product_factory: Callable[..., Product] + ): products = [product_factory(), product_factory(), product_factory()] cnt = len(products) res = product_manager.get_by_uuids(product_uuids=[p.id for p in products]) @@ -43,7 +51,7 @@ class TestProductManagerGetMethods: assert "invalid uuid passed" in str(cm.value) def test_get_by_uuid_if_exists( - self, product_factory: Callable[..., Product], product_manager + self, product_factory: Callable[..., Product], product_manager: ProductManager ): products = [product_factory(), product_factory(), product_factory()] @@ -54,7 +62,7 @@ class TestProductManagerGetMethods: assert instance is None def test_get_by_uuids_if_exists( - self, product_manager: ProductManager, product_factory + self, product_manager: ProductManager, product_factory: Callable[..., Product] ): products = [product_factory(), product_factory(), product_factory()] @@ -78,7 +86,7 @@ class TestProductManagerGetMethods: class TestProductManagerGetAll: @pytest.mark.skip(reason="TODO") - def test_get_ALL_by_ids(self, product_manager): + def test_get_ALL_by_ids(self, product_manager: ProductManager): products = product_manager.get_all(rand_limit=50) logger.info(f"Fetching {len(products)} product uuids") # todo: once timebucks stops spamming broken accounts, fetch more diff --git a/tests/managers/thl/test_profiling/test_question.py b/tests/managers/thl/test_profiling/test_question.py index 998466e..97e7365 100644 --- a/tests/managers/thl/test_profiling/test_question.py +++ b/tests/managers/thl/test_profiling/test_question.py @@ -1,3 +1,4 @@ +from collections.abc import Callable from uuid import uuid4 from generalresearch.managers.thl.profiling.question import QuestionManager @@ -6,7 +7,11 @@ from generalresearch.models import Source class TestQuestionManager: - def test_get_multi_upk(self, question_manager: QuestionManager, upk_data): + def test_get_multi_upk( + self, question_manager: QuestionManager, upk_data: Callable[..., None] + ): + upk_data() + qs = question_manager.get_multi_upk( question_ids=[ "8a22de34f985476aac85e15547100db8", @@ -17,13 +22,21 @@ class TestQuestionManager: ) assert len(qs) == 3 - def test_get_questions_ranked(self, question_manager: QuestionManager, upk_data): + def test_get_questions_ranked( + self, question_manager: QuestionManager, upk_data: Callable[..., None] + ): + upk_data() + qs = question_manager.get_questions_ranked(country_iso="mx", language_iso="spa") assert len(qs) >= 40 assert qs[0].importance.task_score > qs[40].importance.task_score assert all(q.country_iso == "mx" and q.language_iso == "spa" for q in qs) - def test_lookup_by_property(self, question_manager: QuestionManager, upk_data): + def test_lookup_by_property( + self, question_manager: QuestionManager, upk_data: Callable[..., None] + ): + upk_data() + q = question_manager.lookup_by_property( property_code="i:industry", country_iso="us", language_iso="eng" ) @@ -38,7 +51,11 @@ class TestQuestionManager: ) assert q.explanation_template - def test_filter_by_property(self, question_manager: QuestionManager, upk_data): + def test_filter_by_property( + self, question_manager: QuestionManager, upk_data: Callable[..., None] + ): + upk_data() + lookup = [ ("i:industry", "us", "eng"), ("i:industry", "mx", "eng"), diff --git a/tests/managers/thl/test_profiling/test_schema.py b/tests/managers/thl/test_profiling/test_schema.py index ae61527..b0eae31 100644 --- a/tests/managers/thl/test_profiling/test_schema.py +++ b/tests/managers/thl/test_profiling/test_schema.py @@ -1,9 +1,18 @@ +from collections.abc import Callable + +from generalresearch.managers.thl.profiling.schema import ( + UpkSchemaManager, +) from generalresearch.models.thl.profiling.upk_property import PropertyType class TestUpkSchemaManager: - def test_get_props_info(self, upk_schema_manager, upk_data): + def test_get_props_info( + self, upk_schema_manager: UpkSchemaManager, upk_data: Callable[..., None] + ): + upk_data() + props = upk_schema_manager.get_props_info() assert ( len(props) == 16955 @@ -35,10 +44,10 @@ class TestUpkSchemaManager: assert age.prop_type == PropertyType.UPK_NUMERICAL assert age.gold_standard - cars = [ + cars = next( x for x in props if x.country_iso == "us" and x.property_label == "household_auto_type" - ][0] + ) assert not cars.gold_standard assert cars.categories[0].label == "Autos & Vehicles" diff --git a/tests/managers/thl/test_profiling/test_uqa.py b/tests/managers/thl/test_profiling/test_uqa.py deleted file mode 100644 index 8b13789..0000000 --- a/tests/managers/thl/test_profiling/test_uqa.py +++ /dev/null @@ -1 +0,0 @@ - diff --git a/tests/managers/thl/test_profiling/test_user_upk.py b/tests/managers/thl/test_profiling/test_user_upk.py index 8b995b1..fa10b67 100644 --- a/tests/managers/thl/test_profiling/test_user_upk.py +++ b/tests/managers/thl/test_profiling/test_user_upk.py @@ -1,6 +1,8 @@ +from collections.abc import Callable from datetime import UTC, datetime from generalresearch.managers.thl.profiling.user_upk import UserUpkManager +from generalresearch.models.thl.user import User now = datetime.now(tz=UTC) base = { @@ -21,11 +23,25 @@ for a in upk_ans_dict: class TestUserUpkManager: - def test_user_upk_empty(self, user_upk_manager: UserUpkManager, upk_data, user): + def test_user_upk_empty( + self, + user_upk_manager: UserUpkManager, + upk_data: Callable[..., None], + user: User, + ): + upk_data() + res = user_upk_manager.get_user_upk_mysql(user_id=user.user_id) assert len(res) == 0 - def test_user_upk(self, user_upk_manager: UserUpkManager, upk_data, user): + def test_user_upk( + self, + user_upk_manager: UserUpkManager, + upk_data: Callable[..., None], + user: User, + ): + upk_data() + for x in upk_ans_dict: x["user_id"] = user.user_id user_upk = user_upk_manager.populate_user_upk_from_dict(upk_ans_dict) diff --git a/tests/managers/thl/test_session_manager.py b/tests/managers/thl/test_session_manager.py index 6c5f820..05a49c1 100644 --- a/tests/managers/thl/test_session_manager.py +++ b/tests/managers/thl/test_session_manager.py @@ -1,22 +1,34 @@ -from datetime import timedelta +from __future__ import annotations + +from collections.abc import Callable +from datetime import datetime, timedelta from decimal import Decimal from uuid import uuid4 from faker import Faker +from generalresearch.managers.thl.session import SessionManager from generalresearch.models import DeviceType +from generalresearch.models.gr.business import Business +from generalresearch.models.gr.team import Team from generalresearch.models.legacy.bucket import Bucket from generalresearch.models.thl.definitions import ( SessionStatusCode2, Status, StatusCode1, ) +from generalresearch.models.thl.product import Product +from generalresearch.models.thl.session import Session +from generalresearch.models.thl.user import User +from generalresearch.pg_helper import PostgresConfig fake = Faker() class TestSessionManager: - def test_create_session(self, session_manager, user, utc_hour_ago): + def test_create_session( + self, session_manager: SessionManager, user: User, utc_hour_ago: datetime + ): bucket = Bucket( loi_min=timedelta(seconds=60), loi_max=timedelta(seconds=120), @@ -39,7 +51,9 @@ class TestSessionManager: s2 = session_manager.get_from_uuid(session_uuid=s1.uuid) assert s1 == s2 - def test_finish_with_status(self, session_manager, user, utc_hour_ago): + def test_finish_with_status( + self, session_manager: SessionManager, user: User, utc_hour_ago: datetime + ): uuid_1 = uuid4().hex session = session_manager.create( started=utc_hour_ago, user=user, uuid_id=uuid_1 @@ -59,7 +73,7 @@ class TestSessionManager: class TestSessionManagerFilter: - def test_base(self, session_manager, user, utc_now): + def test_base(self, session_manager: SessionManager, user: User, utc_now: datetime): uuid_id = uuid4().hex session_manager.create(started=utc_now, user=user, uuid_id=uuid_id) res = session_manager.filter(limit=1) @@ -67,7 +81,9 @@ class TestSessionManagerFilter: assert isinstance(res, list) assert res[0].uuid == uuid_id - def test_user(self, session_manager, user, utc_hour_ago): + def test_user( + self, session_manager: SessionManager, user: User, utc_hour_ago: datetime + ): session_manager.create(started=utc_hour_ago, user=user, uuid_id=uuid4().hex) session_manager.create(started=utc_hour_ago, user=user, uuid_id=uuid4().hex) @@ -78,16 +94,13 @@ class TestSessionManagerFilter: self, product_factory: Callable[..., Product], user_factory: Callable[..., User], - session_manager, - user, - utc_hour_ago, + session_manager: SessionManager, + utc_hour_ago: datetime, ): - from generalresearch.models.thl.session import Session - from generalresearch.models.thl.user import User p1 = product_factory() - for n in range(5): + for _ in range(5): u = user_factory(product=p1) session_manager.create(started=utc_hour_ago, user=u, uuid_id=uuid4().hex) @@ -102,15 +115,14 @@ class TestSessionManagerFilter: self, product_factory: Callable[..., Product], user_factory: Callable[..., User], - team, - session_manager, - user, - utc_hour_ago, + team: Team, + session_manager: SessionManager, + utc_hour_ago: datetime, thl_web_rr: PostgresConfig, ): p1 = product_factory(team=team) - for n in range(5): + for _ in range(5): u = user_factory(product=p1) session_manager.create(started=utc_hour_ago, user=u, uuid_id=uuid4().hex) @@ -124,14 +136,13 @@ class TestSessionManagerFilter: product_factory: Callable[..., Product], business: Business, user_factory: Callable[..., User], - session_manager, - user, - utc_hour_ago, + session_manager: SessionManager, + utc_hour_ago: datetime, thl_web_rr: PostgresConfig, ): p1 = product_factory(business=business) - for n in range(5): + for _ in range(5): u = user_factory(product=p1) session_manager.create(started=utc_hour_ago, user=u, uuid_id=uuid4().hex) diff --git a/tests/managers/thl/test_survey.py b/tests/managers/thl/test_survey.py index eec7c3b..2c2bf9d 100644 --- a/tests/managers/thl/test_survey.py +++ b/tests/managers/thl/test_survey.py @@ -1,9 +1,21 @@ +from __future__ import annotations + import uuid +from collections.abc import Callable from datetime import UTC, datetime from decimal import Decimal import pytest +from generalresearch.managers.thl.buyer import BuyerManager +from generalresearch.managers.thl.profiling.question import ( + QuestionManager, +) +from generalresearch.managers.thl.profiling.schema import ( + UpkSchemaManager, +) +from generalresearch.managers.thl.profiling.uqa import UQAManager +from generalresearch.managers.thl.survey import SurveyManager, SurveyStatManager from generalresearch.models import Source from generalresearch.models.legacy.bucket import ( DurationSummary, @@ -23,7 +35,7 @@ from generalresearch.models.thl.survey.model import ( @pytest.fixture(scope="session") -def surveys_fixture(): +def surveys_fixture() -> list[Survey]: return [ Survey(source=Source.TESTING, survey_id="a", buyer_code="buyer1"), Survey(source=Source.TESTING, survey_id="b", buyer_code="buyer2"), @@ -73,11 +85,13 @@ class TestSurvey: def test( self, - delete_buyers_surveys, - buyer_manager, - survey_manager, - surveys_fixture, + delete_buyers_surveys: Callable[..., None], + buyer_manager: BuyerManager, + survey_manager: SurveyManager, + surveys_fixture: list[Survey], ): + delete_buyers_surveys() + survey_manager.create_or_update(surveys_fixture) survey_ids = {s.survey_id for s in surveys_fixture} res = survey_manager.filter_by_natural_key( @@ -98,7 +112,7 @@ class TestSurvey: assert res2[0] == res[0] assert len(res2) == len(surveys2) - def test_category(self, survey_manager): + def test_category(self, survey_manager: SurveyManager): survey1 = Survey(id=562289, survey_id="a", source=Source.TESTING) survey2 = Survey(id=562290, survey_id="a", source=Source.TESTING) categories = list(survey_manager.category_manager.categories.values()) @@ -110,8 +124,14 @@ class TestSurvey: survey_manager.update_surveys_categories(surveys) def test_survey_eligibility( - self, survey_manager, upk_data, question_manager, uqa_manager + self, + survey_manager: SurveyManager, + upk_data: Callable[..., None], + question_manager: QuestionManager, + uqa_manager: UQAManager, ): + upk_data() + bucket = TopNPlusBucket( id="c82cf98c578a43218334544ab376b00e", contents=[], @@ -205,10 +225,10 @@ class TestSurvey: class TestSurveyStat: def test( self, - delete_buyers_surveys, + delete_buyers_surveys: Callable[..., None], surveystat_manager, - survey_manager, - surveys_fixture, + survey_manager: SurveyManager, + surveys_fixture: list[Survey], ): survey_manager.create_or_update(surveys_fixture) ss = [ssa, ssb] @@ -276,15 +296,17 @@ class TestSurveyStat: def test_ymsp( self, - delete_buyers_surveys, - surveys_fixture, - survey_manager, - surveystat_manager, + delete_buyers_surveys: Callable[..., None], + surveys_fixture: list[Survey], + survey_manager: SurveyManager, + surveystat_manager: SurveyStatManager, ): + delete_buyers_surveys() + source = Source.TESTING survey = surveys_fixture[0].model_copy() surveys = [] - for idx in range(100): + for _ in range(100): s = survey.model_copy() s.survey_id = uuid.uuid4().hex surveys.append(s) @@ -305,7 +327,7 @@ class TestSurveyStat: surveys = surveys[10:] # and 2 new ones are created - for idx in range(2): + for _ in range(2): s = survey.model_copy() s.survey_id = uuid.uuid4().hex surveys.append(s) @@ -329,11 +351,13 @@ class TestSurveyStat: def test_filter( self, - delete_buyers_surveys, - surveys_fixture, - survey_manager, - surveystat_manager, + delete_buyers_surveys: Callable[..., None], + surveys_fixture: list[Survey], + survey_manager: SurveyManager, + surveystat_manager: SurveyStatManager, ): + delete_buyers_surveys() + surveys = [] survey = surveys_fixture[0].model_copy() survey.source = Source.TESTING diff --git a/tests/managers/thl/test_survey_penalty.py b/tests/managers/thl/test_survey_penalty.py index c7862bb..9c29a0a 100644 --- a/tests/managers/thl/test_survey_penalty.py +++ b/tests/managers/thl/test_survey_penalty.py @@ -1,7 +1,10 @@ +from __future__ import annotations + import uuid import pytest +from generalresearch.managers.thl.survey_penalty import SurveyPenaltyManager from generalresearch.models import Source from generalresearch.models.thl.survey.penalty import ( BPSurveyPenalty, @@ -22,7 +25,9 @@ def team_uuid() -> str: @pytest.fixture -def penalties(product_uuid, team_uuid): +def penalties( + product_uuid: str, team_uuid: str +) -> list[BPSurveyPenalty | TeamSurveyPenalty]: return [ BPSurveyPenalty( source=Source.TESTING, survey_id="a", penalty=0.1, product_id=product_uuid @@ -48,7 +53,13 @@ def penalties(product_uuid, team_uuid): class TestSurveyPenalty: - def test(self, surveypenalty_manager, penalties, product_uuid, team_uuid): + def test( + self, + surveypenalty_manager: SurveyPenaltyManager, + penalties: list[BPSurveyPenalty | TeamSurveyPenalty], + product_uuid: str, + team_uuid: str, + ): surveypenalty_manager.set_penalties(penalties) res = surveypenalty_manager.get_penalties_for( diff --git a/tests/managers/thl/test_task_adjustment.py b/tests/managers/thl/test_task_adjustment.py index 7b77a68..a7324c3 100644 --- a/tests/managers/thl/test_task_adjustment.py +++ b/tests/managers/thl/test_task_adjustment.py @@ -1,20 +1,31 @@ +from __future__ import annotations + import logging +from collections.abc import Callable from datetime import UTC, datetime, timedelta from decimal import Decimal from random import randint import pytest +from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager +from generalresearch.managers.thl.session import SessionManager +from generalresearch.managers.thl.task_adjustment import ( + TaskAdjustmentManager, +) +from generalresearch.managers.thl.wall import WallManager from generalresearch.models import Source from generalresearch.models.thl.definitions import ( Status, StatusCode1, WallAdjustedStatus, ) +from generalresearch.models.thl.session import Session +from generalresearch.models.thl.user import User @pytest.fixture() -def session_complete(session_with_tx_factory: Callable[..., None], user): +def session_complete(session_with_tx_factory: Callable[..., Session], user: User): return session_with_tx_factory( user=user, final_status=Status.COMPLETE, wall_req_cpi=Decimal("1.23") ) @@ -22,7 +33,7 @@ def session_complete(session_with_tx_factory: Callable[..., None], user): @pytest.fixture() def session_complete_with_wallet( - session_with_tx_factory: Callable[..., None], user_with_wallet + session_with_tx_factory: Callable[..., None], user_with_wallet: User ): return session_with_tx_factory( user=user_with_wallet, @@ -32,7 +43,9 @@ def session_complete_with_wallet( @pytest.fixture() -def session_fail(user, session_manager, wall_manager): +def session_fail( + user: User, session_manager: SessionManager, wall_manager: WallManager +) -> Session: session = session_manager.create_dummy(started=datetime.now(UTC), user=user) wall1 = wall_manager.create_dummy( session_id=session.id, @@ -56,38 +69,37 @@ class TestHandleRecons: def test_complete_to_recon( self, - session_complete, - thl_lm, - task_adjustment_manager, - wall_manager, - session_manager, + session_complete: Session, + thl_ledger_manager: ThlLedgerManager, + task_adjustment_manager: TaskAdjustmentManager, + wall_manager: WallManager, + session_manager: SessionManager, caplog, ): print(wall_manager.pg_config.dsn) mid = session_complete.uuid wall_uuid = session_complete.wall_events[-1].uuid s = session_complete - ledger_manager = thl_lm - revenue_account = ledger_manager.get_account_task_complete_revenue() - current_amount = ledger_manager.get_account_filtered_balance( + revenue_account = thl_ledger_manager.get_account_task_complete_revenue() + current_amount = thl_ledger_manager.get_account_filtered_balance( revenue_account, "thl_wall", wall_uuid ) assert ( current_amount == 123 ), "this is the amount of revenue from this task complete" - bp_wallet_account = ledger_manager.get_account_or_create_bp_wallet( + bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet( s.user.product ) - current_bp_payout = ledger_manager.get_account_filtered_balance( + current_bp_payout = thl_ledger_manager.get_account_filtered_balance( bp_wallet_account, "thl_session", mid ) assert current_bp_payout == 117, "this is the amount paid to the BP" # Do the work here !! ----v task_adjustment_manager.handle_single_recon( - ledger_manager=thl_lm, + ledger_manager=thl_ledger_manager, wall_uuid=wall_uuid, adjusted_status=WallAdjustedStatus.ADJUSTED_TO_FAIL, ) @@ -95,18 +107,18 @@ class TestHandleRecons: len(task_adjustment_manager.filter_by_wall_uuid(wall_uuid=wall_uuid)) == 1 ) - current_amount = ledger_manager.get_account_filtered_balance( + current_amount = thl_ledger_manager.get_account_filtered_balance( revenue_account, "thl_wall", wall_uuid ) assert current_amount == 0, "after recon, it should be zeroed" - current_bp_payout = ledger_manager.get_account_filtered_balance( + current_bp_payout = thl_ledger_manager.get_account_filtered_balance( bp_wallet_account, "thl_session", mid ) assert current_bp_payout == 0, "this is the amount paid to the BP" - commission_account = ledger_manager.get_account_or_create_bp_commission( + commission_account = thl_ledger_manager.get_account_or_create_bp_commission( s.user.product ) - assert ledger_manager.get_account_balance(commission_account) == 0 + assert thl_ledger_manager.get_account_balance(commission_account) == 0 # Now, say we get the exact same *adjust to incomplete* msg again. It should do nothing! adjusted_timestamp = datetime.now(tz=UTC) @@ -122,7 +134,7 @@ class TestHandleRecons: session = session_manager.get_from_id(wall.session_id) user = session.user with caplog.at_level(logging.INFO): - ledger_manager.create_tx_task_adjustment( + thl_ledger_manager.create_tx_task_adjustment( wall, user=user, created=adjusted_timestamp ) assert "No transactions needed" in caplog.text @@ -135,212 +147,219 @@ class TestHandleRecons: assert "is already f" in caplog.text or "is already Status.FAIL" in caplog.text with caplog.at_level(logging.INFO, logger="LedgerManager"): - ledger_manager.create_tx_bp_adjustment(session, created=adjusted_timestamp) + thl_ledger_manager.create_tx_bp_adjustment( + session, created=adjusted_timestamp + ) assert "No transactions needed" in caplog.text - current_amount = ledger_manager.get_account_filtered_balance( + current_amount = thl_ledger_manager.get_account_filtered_balance( revenue_account, "thl_wall", wall_uuid ) assert current_amount == 0, "after recon, it should be zeroed" - current_bp_payout = ledger_manager.get_account_filtered_balance( + current_bp_payout = thl_ledger_manager.get_account_filtered_balance( bp_wallet_account, "thl_session", mid ) assert current_bp_payout == 0, "this is the amount paid to the BP" # And if we get an adj to fail, and handle it, it should do nothing at all task_adjustment_manager.handle_single_recon( - ledger_manager=thl_lm, + ledger_manager=thl_ledger_manager, wall_uuid=wall_uuid, adjusted_status=WallAdjustedStatus.ADJUSTED_TO_FAIL, ) assert ( len(task_adjustment_manager.filter_by_wall_uuid(wall_uuid=wall_uuid)) == 1 ) - current_amount = ledger_manager.get_account_filtered_balance( + current_amount = thl_ledger_manager.get_account_filtered_balance( revenue_account, "thl_wall", wall_uuid ) assert current_amount == 0, "after recon, it should be zeroed" - def test_fail_to_complete(self, session_fail, thl_lm, task_adjustment_manager): - s = session_fail + def test_fail_to_complete( + self, + session_fail: Session, + thl_ledger_manager: ThlLedgerManager, + task_adjustment_manager: TaskAdjustmentManager, + ): mid = session_fail.uuid wall_uuid = session_fail.wall_events[-1].uuid - ledger_manager = thl_lm - revenue_account = ledger_manager.get_account_task_complete_revenue() - current_amount = ledger_manager.get_account_filtered_balance( + revenue_account = thl_ledger_manager.get_account_task_complete_revenue() + current_amount = thl_ledger_manager.get_account_filtered_balance( revenue_account, "thl_wall", mid ) assert ( current_amount == 0 ), "this is the amount of revenue from this task complete" - bp_wallet_account = ledger_manager.get_account_or_create_bp_wallet( - s.user.product + bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet( + session_fail.user.product ) - current_bp_payout = ledger_manager.get_account_filtered_balance( + current_bp_payout = thl_ledger_manager.get_account_filtered_balance( bp_wallet_account, "thl_session", mid ) assert current_bp_payout == 0, "this is the amount paid to the BP" task_adjustment_manager.handle_single_recon( - ledger_manager=thl_lm, + ledger_manager=thl_ledger_manager, wall_uuid=wall_uuid, adjusted_status=WallAdjustedStatus.ADJUSTED_TO_COMPLETE, ) - current_amount = ledger_manager.get_account_filtered_balance( + current_amount = thl_ledger_manager.get_account_filtered_balance( revenue_account, "thl_wall", wall_uuid ) assert current_amount == 322, "after recon, we should be paid" - current_bp_payout = ledger_manager.get_account_filtered_balance( + current_bp_payout = thl_ledger_manager.get_account_filtered_balance( bp_wallet_account, "thl_session", mid ) assert current_bp_payout == 306, "this is the amount paid to the BP" # Now reverse it back to fail task_adjustment_manager.handle_single_recon( - ledger_manager=thl_lm, + ledger_manager=thl_ledger_manager, wall_uuid=wall_uuid, adjusted_status=WallAdjustedStatus.ADJUSTED_TO_FAIL, ) - current_amount = ledger_manager.get_account_filtered_balance( + current_amount = thl_ledger_manager.get_account_filtered_balance( revenue_account, "thl_wall", wall_uuid ) assert current_amount == 0 - current_bp_payout = ledger_manager.get_account_filtered_balance( + current_bp_payout = thl_ledger_manager.get_account_filtered_balance( bp_wallet_account, "thl_session", mid ) assert current_bp_payout == 0 - commission_account = ledger_manager.get_account_or_create_bp_commission( - s.user.product + commission_account = thl_ledger_manager.get_account_or_create_bp_commission( + session_fail.user.product ) - assert ledger_manager.get_account_balance(commission_account) == 0 + assert thl_ledger_manager.get_account_balance(commission_account) == 0 def test_complete_already_complete( - self, session_complete, thl_lm, task_adjustment_manager + self, + session_complete: Session, + thl_ledger_manager: ThlLedgerManager, + task_adjustment_manager: TaskAdjustmentManager, ): - s = session_complete mid = session_complete.uuid wall_uuid = session_complete.wall_events[-1].uuid - ledger_manager = thl_lm for _ in range(4): # just run it 4 times to make sure nothing happens 4 times task_adjustment_manager.handle_single_recon( - ledger_manager=thl_lm, + ledger_manager=thl_ledger_manager, wall_uuid=wall_uuid, adjusted_status=WallAdjustedStatus.ADJUSTED_TO_COMPLETE, ) - revenue_account = ledger_manager.get_account_task_complete_revenue() - bp_wallet_account = ledger_manager.get_account_or_create_bp_wallet( - s.user.product + revenue_account = thl_ledger_manager.get_account_task_complete_revenue() + bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet( + session_complete.user.product ) - commission_account = ledger_manager.get_account_or_create_bp_commission( - s.user.product + commission_account = thl_ledger_manager.get_account_or_create_bp_commission( + session_complete.user.product ) - current_amount = ledger_manager.get_account_filtered_balance( + current_amount = thl_ledger_manager.get_account_filtered_balance( revenue_account, "thl_wall", wall_uuid ) assert current_amount == 123 - assert ledger_manager.get_account_balance(commission_account) == 6 + assert thl_ledger_manager.get_account_balance(commission_account) == 6 - current_bp_payout = ledger_manager.get_account_filtered_balance( + current_bp_payout = thl_ledger_manager.get_account_filtered_balance( bp_wallet_account, "thl_session", mid ) assert current_bp_payout == 117 def test_incomplete_already_incomplete( - self, session_fail, thl_lm, task_adjustment_manager + self, + session_fail: Session, + thl_ledger_manager: ThlLedgerManager, + task_adjustment_manager: TaskAdjustmentManager, ): - s = session_fail mid = session_fail.uuid wall_uuid = session_fail.wall_events[-1].uuid - ledger_manager = thl_lm for _ in range(4): task_adjustment_manager.handle_single_recon( - ledger_manager=thl_lm, + ledger_manager=thl_ledger_manager, wall_uuid=wall_uuid, adjusted_status=WallAdjustedStatus.ADJUSTED_TO_FAIL, ) - revenue_account = ledger_manager.get_account_task_complete_revenue() - bp_wallet_account = ledger_manager.get_account_or_create_bp_wallet( - s.user.product + revenue_account = thl_ledger_manager.get_account_task_complete_revenue() + bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet( + session_fail.user.product ) - commission_account = ledger_manager.get_account_or_create_bp_commission( - s.user.product + commission_account = thl_ledger_manager.get_account_or_create_bp_commission( + session_fail.user.product ) - current_amount = ledger_manager.get_account_filtered_balance( + current_amount = thl_ledger_manager.get_account_filtered_balance( revenue_account, "thl_wall", mid ) assert current_amount == 0 - assert ledger_manager.get_account_balance(commission_account) == 0 + assert thl_ledger_manager.get_account_balance(commission_account) == 0 - current_bp_payout = ledger_manager.get_account_filtered_balance( + current_bp_payout = thl_ledger_manager.get_account_filtered_balance( bp_wallet_account, "thl_session", mid ) assert current_bp_payout == 0 def test_complete_to_recon_user_wallet( self, - session_complete_with_wallet, - user_with_wallet, - thl_lm, - task_adjustment_manager, + session_complete_with_wallet: Session, + # user_with_wallet: User, + thl_ledger_manager: ThlLedgerManager, + task_adjustment_manager: TaskAdjustmentManager, ): - s = session_complete_with_wallet - mid = s.uuid - wall_uuid = s.wall_events[-1].uuid - ledger_manager = thl_lm + mid = session_complete_with_wallet.uuid + wall_uuid = session_complete_with_wallet.wall_events[-1].uuid - revenue_account = ledger_manager.get_account_task_complete_revenue() - amount = ledger_manager.get_account_filtered_balance( + revenue_account = thl_ledger_manager.get_account_task_complete_revenue() + amount = thl_ledger_manager.get_account_filtered_balance( revenue_account, "thl_wall", wall_uuid ) assert amount == 123, "this is the amount of revenue from this task complete" - bp_wallet_account = ledger_manager.get_account_or_create_bp_wallet( - s.user.product + bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet( + session_complete_with_wallet.user.product ) - user_wallet_account = ledger_manager.get_account_or_create_user_wallet(s.user) - commission_account = ledger_manager.get_account_or_create_bp_commission( - s.user.product + user_wallet_account = thl_ledger_manager.get_account_or_create_user_wallet( + session_complete_with_wallet.user + ) + commission_account = thl_ledger_manager.get_account_or_create_bp_commission( + session_complete_with_wallet.user.product ) - amount = ledger_manager.get_account_filtered_balance( + amount = thl_ledger_manager.get_account_filtered_balance( bp_wallet_account, "thl_session", mid ) assert amount == 70, "this is the amount paid to the BP" - amount = ledger_manager.get_account_filtered_balance( + amount = thl_ledger_manager.get_account_filtered_balance( user_wallet_account, "thl_session", mid ) assert amount == 47, "this is the amount paid to the user" assert ( - ledger_manager.get_account_balance(commission_account) == 6 + thl_ledger_manager.get_account_balance(commission_account) == 6 ), "earned commission" task_adjustment_manager.handle_single_recon( - ledger_manager=thl_lm, + ledger_manager=thl_ledger_manager, wall_uuid=wall_uuid, adjusted_status=WallAdjustedStatus.ADJUSTED_TO_FAIL, ) - amount = ledger_manager.get_account_filtered_balance( + amount = thl_ledger_manager.get_account_filtered_balance( revenue_account, "thl_wall", wall_uuid ) assert amount == 0 - amount = ledger_manager.get_account_filtered_balance( + amount = thl_ledger_manager.get_account_filtered_balance( bp_wallet_account, "thl_session", mid ) assert amount == 0 - amount = ledger_manager.get_account_filtered_balance( + amount = thl_ledger_manager.get_account_filtered_balance( user_wallet_account, "thl_session", mid ) assert amount == 0 assert ( - ledger_manager.get_account_balance(commission_account) == 0 + thl_ledger_manager.get_account_balance(commission_account) == 0 ), "earned commission" diff --git a/tests/managers/thl/test_task_status.py b/tests/managers/thl/test_task_status.py index 44938a6..b47f650 100644 --- a/tests/managers/thl/test_task_status.py +++ b/tests/managers/thl/test_task_status.py @@ -1,9 +1,14 @@ +from __future__ import annotations + +from collections.abc import Callable from datetime import UTC, datetime, timedelta from decimal import Decimal import pytest +from generalresearch.managers.thl.product import ProductManager from generalresearch.managers.thl.session import SessionManager +from generalresearch.managers.thl.wall import WallManager from generalresearch.models import Source from generalresearch.models.thl.definitions import ( Status, @@ -14,6 +19,7 @@ from generalresearch.models.thl.product import ( PayoutConfig, PayoutTransformation, PayoutTransformationPercentArgs, + Product, UserWalletConfig, ) from generalresearch.models.thl.session import Session, WallOut @@ -30,7 +36,7 @@ finish3 = start3 + timedelta(minutes=5) @pytest.fixture(scope="session") -def bp1(product_manager): +def bp1(product_manager: ProductManager) -> Product: # user wallet disabled, payout xform NULL return product_manager.create_dummy( user_wallet_config=UserWalletConfig(enabled=False), @@ -39,7 +45,7 @@ def bp1(product_manager): @pytest.fixture(scope="session") -def bp2(product_manager): +def bp2(product_manager: ProductManager) -> Product: # user wallet disabled, payout xform 40% return product_manager.create_dummy( user_wallet_config=UserWalletConfig(enabled=False), @@ -53,7 +59,7 @@ def bp2(product_manager): @pytest.fixture(scope="session") -def bp3(product_manager): +def bp3(product_manager: ProductManager) -> Product: # user wallet enabled, payout xform 50% return product_manager.create_dummy( user_wallet_config=UserWalletConfig(enabled=True), @@ -70,9 +76,9 @@ class TestTaskStatus: def test_task_status_complete_1( self, - bp1, + bp1: Product, user_factory: Callable[..., User], - finished_session_factory, + finished_session_factory: Callable[..., Session], session_manager: SessionManager, ): # User Payout xform NULL @@ -131,10 +137,10 @@ class TestTaskStatus: def test_task_status_complete_2( self, - bp2, + bp2: Product, user_factory: Callable[..., User], - finished_session_factory, - session_manager, + finished_session_factory: Callable[..., Session], + session_manager: SessionManager, ): # User Payout xform 40% user2: User = user_factory(product=bp2) @@ -202,10 +208,10 @@ class TestTaskStatus: def test_task_status_complete_3( self, - bp3, + bp3: Product, user_factory: Callable[..., User], - finished_session_factory, - session_manager, + finished_session_factory: Callable[..., Session], + session_manager: SessionManager, ): # Wallet enabled User Payout xform 50% (the response is identical # to the user wallet disabled w same xform) @@ -235,16 +241,17 @@ class TestTaskStatus: } ) tsr = session_manager.get_task_status_response(s3.uuid) + assert isinstance(tsr, TaskStatusResponse) # Not bothering with wall events ... tsr.wall_events = None assert tsr == expected_tsr def test_task_status_fail( self, - bp1, + bp1: Product, user_factory: Callable[..., User], - finished_session_factory, - session_manager, + finished_session_factory: Callable[..., Session], + session_manager: SessionManager, ): # User Payout xform NULL: user payout is None always user1: User = user_factory(product=bp1) @@ -275,16 +282,17 @@ class TestTaskStatus: } ) tsr = session_manager.get_task_status_response(s1.uuid) + assert isinstance(tsr, TaskStatusResponse) # Not bothering with wall events ... tsr.wall_events = None assert tsr == expected_tsr def test_task_status_fail_xform( self, - bp2, + bp2: Product, user_factory: Callable[..., User], - finished_session_factory, - session_manager, + finished_session_factory: Callable[..., Session], + session_manager: SessionManager, ): # User Payout xform 40%: user_payout is 0 (not None) @@ -314,16 +322,17 @@ class TestTaskStatus: } ) tsr = session_manager.get_task_status_response(s.uuid) + assert isinstance(tsr, TaskStatusResponse) # Not bothering with wall events ... tsr.wall_events = None assert tsr == expected_tsr def test_task_status_abandon( self, - bp1, + bp1: Product, user_factory: Callable[..., User], - session_factory, - session_manager, + session_factory: Callable[..., Session], + session_manager: SessionManager, ): # User Payout xform NULL: all payout fields are None user: User = user_factory(product=bp1) @@ -352,16 +361,17 @@ class TestTaskStatus: } ) tsr = session_manager.get_task_status_response(s.uuid) + assert isinstance(tsr, TaskStatusResponse) # Not bothering with wall events ... tsr.wall_events = None assert tsr == expected_tsr def test_task_status_abandon_xform( self, - bp2, + bp2: Product, user_factory: Callable[..., User], - session_factory, - session_manager, + session_factory: Callable[..., Session], + session_manager: SessionManager, ): # User Payout xform 40%: all payout fields are None (same as when payout xform is null) user: User = user_factory(product=bp2) @@ -393,17 +403,18 @@ class TestTaskStatus: } ) tsr = session_manager.get_task_status_response(s.uuid) + assert isinstance(tsr, TaskStatusResponse) # Not bothering with wall events ... tsr.wall_events = None assert tsr == expected_tsr def test_task_status_adj_fail( self, - bp1, + bp1: Product, user_factory: Callable[..., User], - finished_session_factory, - wall_manager, - session_manager, + finished_session_factory: Callable[..., Session], + wall_manager: WallManager, + session_manager: SessionManager, ): # Complete -> Fail # User Payout xform NULL: adjusted_user_* and user_* is still all None @@ -442,17 +453,18 @@ class TestTaskStatus: } ) tsr = session_manager.get_task_status_response(s.uuid) + assert isinstance(tsr, TaskStatusResponse) # Not bothering with wall events ... tsr.wall_events = None assert tsr == expected_tsr def test_task_status_adj_fail_xform( self, - bp2, + bp2: Product, user_factory: Callable[..., User], - finished_session_factory, - wall_manager, - session_manager, + finished_session_factory: Callable[..., Session], + wall_manager: WallManager, + session_manager: SessionManager, ): # Complete -> Fail # User Payout xform 40%: adjusted_user_payout is 0 (not null) @@ -494,17 +506,18 @@ class TestTaskStatus: } ) tsr = session_manager.get_task_status_response(s.uuid) + assert isinstance(tsr, TaskStatusResponse) # Not bothering with wall events ... tsr.wall_events = None assert tsr == expected_tsr def test_task_status_adj_complete_from_abandon( self, - bp1, + bp1: Product, user_factory: Callable[..., User], - session_factory, - wall_manager, - session_manager, + session_factory: Callable[..., Session], + wall_manager: WallManager, + session_manager: SessionManager, ): # User Payout xform NULL user: User = user_factory(product=bp1) @@ -548,17 +561,18 @@ class TestTaskStatus: } ) tsr = session_manager.get_task_status_response(s.uuid) + assert isinstance(tsr, TaskStatusResponse) # Not bothering with wall events ... tsr.wall_events = None assert tsr == expected_tsr def test_task_status_adj_complete_from_abandon_xform( self, - bp2, + bp2: Product, user_factory: Callable[..., User], - session_factory, - wall_manager, - session_manager, + session_factory: Callable[..., Session], + wall_manager: WallManager, + session_manager: SessionManager, ): # User Payout xform 40% user: User = user_factory(product=bp2) @@ -605,17 +619,18 @@ class TestTaskStatus: } ) tsr = session_manager.get_task_status_response(s.uuid) + assert isinstance(tsr, TaskStatusResponse) # Not bothering with wall events ... tsr.wall_events = None assert tsr == expected_tsr def test_task_status_adj_complete_from_fail( self, - bp1, + bp1: Product, user_factory: Callable[..., User], - finished_session_factory, - wall_manager, - session_manager, + finished_session_factory: Callable[..., Session], + wall_manager: WallManager, + session_manager: SessionManager, ): # User Payout xform NULL user: User = user_factory(product=bp1) @@ -659,17 +674,18 @@ class TestTaskStatus: } ) tsr = session_manager.get_task_status_response(s.uuid) + assert isinstance(tsr, TaskStatusResponse) # Not bothering with wall events ... tsr.wall_events = None assert tsr == expected_tsr def test_task_status_adj_complete_from_fail_xform( self, - bp2, + bp2: Product, user_factory: Callable[..., User], - finished_session_factory, - wall_manager, - session_manager, + finished_session_factory: Callable[..., Session], + wall_manager: WallManager, + session_manager: SessionManager, ): # User Payout xform 40% user: User = user_factory(product=bp2) @@ -715,6 +731,7 @@ class TestTaskStatus: } ) tsr = session_manager.get_task_status_response(s.uuid) + assert isinstance(tsr, TaskStatusResponse) # Not bothering with wall events ... tsr.wall_events = None assert tsr == expected_tsr diff --git a/tests/managers/thl/test_user_manager/test_base.py b/tests/managers/thl/test_user_manager/test_base.py index 6b259ff..5822207 100644 --- a/tests/managers/thl/test_user_manager/test_base.py +++ b/tests/managers/thl/test_user_manager/test_base.py @@ -10,14 +10,18 @@ from generalresearch.managers.thl.user_manager import ( UserCreateNotAllowedError, get_bp_user_create_limit_hourly, ) +from generalresearch.managers.thl.user_manager.mysql_user_manager import ( + MysqlUserManager, +) from generalresearch.managers.thl.user_manager.rate_limit import ( RateLimitItemPerHourConstantKey, + UserManagerLimiter, ) from generalresearch.managers.thl.user_manager.user_manager import ( UserManager, ) from generalresearch.managers.thl.userhealth import AuditLogManager -from generalresearch.models.thl.product import product: Product, UserCreateConfig +from generalresearch.models.thl.product import Product, UserCreateConfig, product from generalresearch.models.thl.user import User from generalresearch.pg_helper import PostgresConfig @@ -86,10 +90,11 @@ class TestUserManager: class TestBlockUserManager: - def test_block_user(self, product: product: Product, user_manager: UserManager): + def test_block_user(self, product: Product, user_manager: UserManager): product_user_id = f"user-{uuid4().hex[:10]}" # mysql_user_manager to skip user creation limit check + assert isinstance(user_manager.mysql_user_manager, MysqlUserManager) user: User = user_manager.mysql_user_manager.create_user( product_id=product.id, product_user_id=product_user_id ) @@ -113,11 +118,12 @@ class TestBlockUserManager: assert user.blocked def test_block_user_whitelist( - self, product: product: Product, user_manager: UserManager, thl_web_rw: PostgresConfig + self, product: Product, user_manager: UserManager, thl_web_rw: PostgresConfig ): product_user_id = f"user-{uuid4().hex[:10]}" # mysql_user_manager to skip user creation limit check + assert isinstance(user_manager.mysql_user_manager, MysqlUserManager) user: User = user_manager.mysql_user_manager.create_user( product_id=product.id, product_user_id=product_user_id ) @@ -154,6 +160,7 @@ class TestCreateUserManager: product_user_id = f"user-{uuid4().hex[:10]}" + assert isinstance(user_manager.mysql_user_manager, MysqlUserManager) user: User = user_manager.mysql_user_manager.create_user( product_id=product.id, product_user_id=product_user_id ) @@ -200,6 +207,7 @@ class TestCreateUserManager: product_user_id = f"user-{uuid4().hex[:10]}" rand_msg = f"log-{uuid4().hex}" + assert isinstance(user_manager.mysql_user_manager, MysqlUserManager) with caplog.at_level(logging.INFO): logger.info(rand_msg) user1 = user_manager.mysql_user_manager.create_user( @@ -264,6 +272,7 @@ class TestCreateUserManager: assert key == f"LIMITER/thl-grpc/allow_user_create/{instance.id}" # make sure we clear the key or subsequent tests will fail + assert isinstance(user_manager.user_manager_limiter, UserManagerLimiter) user_manager.user_manager_limiter.storage.clear(key=key) n = 0 diff --git a/tests/managers/thl/test_user_manager/test_mysql.py b/tests/managers/thl/test_user_manager/test_mysql.py index d414a13..e6f43ef 100644 --- a/tests/managers/thl/test_user_manager/test_mysql.py +++ b/tests/managers/thl/test_user_manager/test_mysql.py @@ -1,24 +1,25 @@ +from __future__ import annotations + +from generalresearch.managers.thl.user_manager.mysql_user_manager import ( + MysqlUserManager, +) +from generalresearch.models.thl.user import User class TestUserManagerMysqlNew: - def test_get_notset(self, user_manager): - assert ( - user_manager.mysql_user_manager.get_user_from_mysql(user_id=-3105) is None - ) + def test_get_notset(self, mysql_user_manager: MysqlUserManager): + assert mysql_user_manager.get_user_from_mysql(user_id=-3105) is None - def test_get_user_id(self, user, user_manager): - assert ( - user_manager.mysql_user_manager.get_user_from_mysql(user_id=user.user_id) - == user - ) + def test_get_user_id(self, user: User, mysql_user_manager: MysqlUserManager): + assert mysql_user_manager.get_user_from_mysql(user_id=user.user_id) == user - def test_get_uuid(self, user, user_manager): - u = user_manager.mysql_user_manager.get_user_from_mysql(user_uuid=user.uuid) + def test_get_uuid(self, user: User, mysql_user_manager: MysqlUserManager): + u = mysql_user_manager.get_user_from_mysql(user_uuid=user.uuid) assert u == user - def test_get_ubp(self, user, user_manager): - u = user_manager.mysql_user_manager.get_user_from_mysql( + def test_get_ubp(self, user: User, mysql_user_manager: MysqlUserManager): + u = mysql_user_manager.get_user_from_mysql( product_id=user.product_id, product_user_id=user.product_user_id ) assert u == user diff --git a/tests/managers/thl/test_user_manager/test_redis.py b/tests/managers/thl/test_user_manager/test_redis.py index 0731438..04071ee 100644 --- a/tests/managers/thl/test_user_manager/test_redis.py +++ b/tests/managers/thl/test_user_manager/test_redis.py @@ -1,29 +1,37 @@ +from __future__ import annotations + import pytest +from generalresearch.config import GRLBaseSettings from generalresearch.managers.base import Permission +from generalresearch.managers.thl.user_manager.redis_user_manager import ( + RedisUserManager, +) +from generalresearch.models.thl.user import User +from generalresearch.pg_helper import PostgresConfig class TestUserManagerRedis: - def test_get_notset(self, user_manager, user): - user_manager.clear_user_inmemory_cache(user=user) - assert user_manager.redis_user_manager.get_user(user_id=user.user_id) is None + def test_get_notset(self, redis_user_manager: RedisUserManager, user: User): + redis_user_manager.clear_user_inmemory_cache(user=user) + assert redis_user_manager.get_user(user_id=user.user_id) is None - def test_get_user_id(self, user_manager, user): - user_manager.redis_user_manager.set_user(user=user) + def test_get_user_id(self, redis_user_manager: RedisUserManager, user: User): + redis_user_manager.set_user(user=user) - assert user_manager.redis_user_manager.get_user(user_id=user.user_id) == user + assert redis_user_manager.get_user(user_id=user.user_id) == user - def test_get_uuid(self, user_manager, user): - user_manager.redis_user_manager.set_user(user=user) + def test_get_uuid(self, redis_user_manager: RedisUserManager, user: User): + redis_user_manager.set_user(user=user) - assert user_manager.redis_user_manager.get_user(user_uuid=user.uuid) == user + assert redis_user_manager.get_user(user_uuid=user.uuid) == user - def test_get_ubp(self, user_manager, user): - user_manager.redis_user_manager.set_user(user=user) + def test_get_ubp(self, redis_user_manager: RedisUserManager, user: User): + redis_user_manager.set_user(user=user) assert ( - user_manager.redis_user_manager.get_user( + redis_user_manager.get_user( product_id=user.product_id, product_user_id=user.product_user_id ) == user @@ -34,7 +42,13 @@ class TestUserManagerRedis: # I mean, the sets are implicitly tested by the get tests above. no point pass - def test_get_with_cache_prefix(self, settings, user, thl_web_rw, thl_web_rr): + def test_get_with_cache_prefix( + self, + settings: GRLBaseSettings, + user: User, + thl_web_rw: PostgresConfig, + thl_web_rr: PostgresConfig, + ): """ Confirm the prefix functionality is working; we do this so it is easier to migrate between any potentially breaking versions @@ -47,7 +61,7 @@ class TestUserManagerRedis: um1 = UserManager( pg_config=thl_web_rw, - pg_config_rr=thl_web_rr: PostgresConfig, + pg_config_rr=thl_web_rr, sql_permissions=[Permission.UPDATE, Permission.CREATE], redis=settings.redis, redis_timeout=settings.redis_timeout, @@ -55,7 +69,7 @@ class TestUserManagerRedis: um2 = UserManager( pg_config=thl_web_rw, - pg_config_rr=thl_web_rr: PostgresConfig, + pg_config_rr=thl_web_rr, sql_permissions=[Permission.UPDATE, Permission.CREATE], redis=settings.redis, redis_timeout=settings.redis_timeout, @@ -69,9 +83,11 @@ class TestUserManagerRedis: product_id=user.product_id, product_user_id=user.product_user_id ) + assert isinstance(um1.redis_user_manager, RedisUserManager) res1 = um1.redis_user_manager.client.get(f"user-lookup:user_id:{user.user_id}") assert res1 is not None + assert isinstance(um2.redis_user_manager, RedisUserManager) res2 = um2.redis_user_manager.client.get( f"user-lookup-v2:user_id:{user.user_id}" ) diff --git a/tests/managers/thl/test_user_manager/test_user_fetch.py b/tests/managers/thl/test_user_manager/test_user_fetch.py index 5c608b3..87d010a 100644 --- a/tests/managers/thl/test_user_manager/test_user_fetch.py +++ b/tests/managers/thl/test_user_manager/test_user_fetch.py @@ -1,14 +1,22 @@ +from __future__ import annotations + +from collections.abc import Callable from uuid import uuid4 import pytest +from generalresearch.managers.thl.user_manager.user_manager import UserManager +from generalresearch.models.thl.product import Product from generalresearch.models.thl.user import User class TestUserManagerFetch: def test_fetch( - self, user_factory: Callable[..., User], product: Product, user_manager + self, + user_factory: Callable[..., User], + product: Product, + user_manager: UserManager, ): user1: User = user_factory(product=product) user2: User = user_factory(product=product) @@ -31,7 +39,7 @@ class TestUserManagerFetch: res = user_manager.fetch(user_uuids=[uuid4().hex]) assert len(res) == 0 - def test_fetch_invalid(self, user_manager): + def test_fetch_invalid(self, user_manager: UserManager): with pytest.raises(AssertionError) as e: user_manager.fetch(user_uuids=[], user_ids=None) assert "Must pass ONE of user_ids, user_uuids" in str(e.value) diff --git a/tests/managers/thl/test_user_manager/test_user_metadata.py b/tests/managers/thl/test_user_manager/test_user_metadata.py index 0b99afe..670e38a 100644 --- a/tests/managers/thl/test_user_manager/test_user_metadata.py +++ b/tests/managers/thl/test_user_manager/test_user_metadata.py @@ -1,21 +1,35 @@ +from __future__ import annotations + +from collections.abc import Callable from uuid import uuid4 import pytest +from generalresearch.managers.thl.user_manager.user_metadata_manager import ( + UserMetadataManager, +) +from generalresearch.models.thl.product import Product +from generalresearch.models.thl.user import User from generalresearch.models.thl.user_profile import UserMetadata class TestUserMetadataManager: - def test_get_notset(self, user, user_manager, user_metadata_manager): + def test_get_notset( + self, + user: User, + user_metadata_manager: UserMetadataManager, + ): # The row in the db won't exist. It just returns the default obj with everything None (except for the user_id) um1 = user_metadata_manager.get(user_id=user.user_id) assert um1 == UserMetadata(user_id=user.user_id) def test_create( - self, user_factory: Callable[..., User], product: Product, user_metadata_manager + self, + user_factory: Callable[..., User], + product: Product, + user_metadata_manager: UserMetadataManager, ): - from generalresearch.models.thl.user import User u1: User = user_factory(product=product) @@ -29,9 +43,11 @@ class TestUserMetadataManager: assert um == um2 def test_create_no_email( - self, product: Product, user_factory: Callable[..., User], user_metadata_manager + self, + product: Product, + user_factory: Callable[..., User], + user_metadata_manager: UserMetadataManager, ): - from generalresearch.models.thl.user import User u1: User = user_factory(product=product) um = UserMetadata(user_id=u1.user_id) @@ -42,9 +58,11 @@ class TestUserMetadataManager: assert um == um2 def test_update( - self, product: Product, user_factory: Callable[..., User], user_metadata_manager + self, + product: Product, + user_factory: Callable[..., User], + user_metadata_manager: UserMetadataManager, ): - from generalresearch.models.thl.user import User u: User = user_factory(product=product) @@ -66,7 +84,6 @@ class TestUserMetadataManager: def test_filter( self, user_factory: Callable[..., User], product: Product, user_metadata_manager ): - from generalresearch.models.thl.user import User user1: User = user_factory(product=product) user2: User = user_factory(product=product) diff --git a/tests/managers/thl/test_user_streak.py b/tests/managers/thl/test_user_streak.py index e87869f..61e2947 100644 --- a/tests/managers/thl/test_user_streak.py +++ b/tests/managers/thl/test_user_streak.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import copy from datetime import UTC, date, datetime, timedelta from decimal import Decimal @@ -5,8 +7,13 @@ from zoneinfo import ZoneInfo import pytest -from generalresearch.managers.thl.user_streak import compute_streaks_from_days +from generalresearch.managers.thl.session import SessionManager +from generalresearch.managers.thl.user_streak import ( + UserStreakManager, + compute_streaks_from_days, +) from generalresearch.models.thl.definitions import Status, StatusCode1 +from generalresearch.models.thl.user import User from generalresearch.models.thl.user_streak import ( StreakFulfillment, StreakPeriod, @@ -59,7 +66,7 @@ def test_compute_streaks_from_days(): @pytest.fixture -def broken_active_streak(user): +def broken_active_streak(user: User) -> list[UserStreak]: return [ UserStreak( period=StreakPeriod.DAY, @@ -94,7 +101,7 @@ def broken_active_streak(user): ] -def create_session_fail(session_manager, start, user): +def create_session_fail(session_manager: SessionManager, start: datetime, user: User): session = session_manager.create_dummy(started=start, country_iso="us", user=user) session_manager.finish_with_status( session, @@ -104,7 +111,9 @@ def create_session_fail(session_manager, start, user): ) -def create_session_complete(session_manager, start, user): +def create_session_complete( + session_manager: SessionManager, start: datetime, user: User +): session = session_manager.create_dummy(started=start, country_iso="us", user=user) session_manager.finish_with_status( session, @@ -115,7 +124,7 @@ def create_session_complete(session_manager, start, user): ) -def test_user_streak_empty(user_streak_manager, user): +def test_user_streak_empty(user_streak_manager: UserStreakManager, user: User): streaks = user_streak_manager.get_user_streaks( user_id=user.user_id, country_iso="us" ) @@ -123,7 +132,10 @@ def test_user_streak_empty(user_streak_manager, user): def test_user_streaks_active_broken( - user_streak_manager, user, session_manager, broken_active_streak + user_streak_manager: UserStreakManager, + user: User, + session_manager: SessionManager, + broken_active_streak: list[UserStreak], ): # Testing active streak, but broken (not today or yesterday) start1 = datetime(2025, 2, 12, tzinfo=UTC) @@ -171,7 +183,9 @@ def test_user_streaks_active_broken( assert streaks == expected_streaks -def test_user_streak_complete_active(user_streak_manager, user, session_manager): +def test_user_streak_complete_active( + user_streak_manager: UserStreakManager, user: User, session_manager: SessionManager +): """Testing active streak that is today""" # They completed yesterday NY time. Today isn't over so streak is pending @@ -217,9 +231,9 @@ def test_user_streak_complete_active(user_streak_manager, user, session_manager) streaks = user_streak_manager.get_user_streaks( user_id=user.user_id, country_iso="us" ) - streak = [ + streak = next( s for s in streaks if s.fulfillment == StreakFulfillment.COMPLETE and s.period == StreakPeriod.DAY - ][0] + ) assert streak == expected_streak diff --git a/tests/managers/thl/test_userhealth.py b/tests/managers/thl/test_userhealth.py index 98d8f25..ea54359 100644 --- a/tests/managers/thl/test_userhealth.py +++ b/tests/managers/thl/test_userhealth.py @@ -1,3 +1,6 @@ +from __future__ import annotations + +from collections.abc import Callable from datetime import UTC, datetime from uuid import uuid4 @@ -5,23 +8,27 @@ import faker import pytest from generalresearch.managers.thl.userhealth import ( + AuditLogManager, IPRecordManager, UserIpHistoryManager, ) -from generalresearch.models.thl.ipinfo import GeoIPInformation +from generalresearch.models.thl.ipinfo import GeoIPInformation, IPGeoname, IPInformation +from generalresearch.models.thl.product import Product +from generalresearch.models.thl.user import User from generalresearch.models.thl.user_iphistory import ( IPRecord, + UserIPHistory, ) from generalresearch.models.thl.userhealth import AuditLog, AuditLogLevel +from generalresearch.pg_helper import PostgresConfig +from generalresearch.redis_helper import RedisConfig fake = faker.Faker() class TestAuditLog: - def test_init(self, thl_web_rr: PostgresConfig, audit_log_manager): - from generalresearch.managers.thl.userhealth import AuditLogManager - + def test_init(self, thl_web_rr: PostgresConfig, audit_log_manager: AuditLogManager): alm = AuditLogManager(pg_config=thl_web_rr) assert isinstance(alm, AuditLogManager) @@ -33,15 +40,16 @@ class TestAuditLog: argnames="level", argvalues=list(AuditLogLevel), ) - def test_create(self, audit_log_manager, user, level): + def test_create( + self, audit_log_manager: AuditLogManager, user: User, level: AuditLogLevel + ): instance = audit_log_manager.create( user_id=user.user_id, level=level, event_type=uuid4().hex ) assert isinstance(instance, AuditLog) assert instance.id != 1 - def test_get_by_id(self, audit_log, audit_log_manager): - from generalresearch.models.thl.userhealth import AuditLog + def test_get_by_id(self, audit_log: AuditLog, audit_log_manager: AuditLogManager): with pytest.raises(expected_exception=Exception) as cm: audit_log_manager.get_by_id(auditlog_id=999_999_999_999) @@ -57,8 +65,8 @@ class TestAuditLog: self, user_factory: Callable[..., User], product_factory: Callable[..., Product], - audit_log_factory, - audit_log_manager, + audit_log_factory: Callable[..., AuditLog], + audit_log_manager: AuditLogManager, ): p1 = product_factory() p2 = product_factory() @@ -82,7 +90,11 @@ class TestAuditLog: assert len(res) == 1 def test_filter_by_user_id( - self, user_factory: Callable[..., User], product: Product, audit_log_factory, audit_log_manager + self, + user_factory: Callable[..., User], + product: Product, + audit_log_factory: Callable[..., AuditLog], + audit_log_manager: AuditLogManager, ): u1 = user_factory(product=product) u2 = user_factory(product=product) @@ -110,8 +122,8 @@ class TestAuditLog: self, user_factory: Callable[..., User], product_factory: Callable[..., Product], - audit_log_factory, - audit_log_manager, + audit_log_factory: Callable[..., AuditLog], + audit_log_manager: AuditLogManager, ): p1 = product_factory() p2 = product_factory() @@ -144,8 +156,8 @@ class TestAuditLog: self, user_factory: Callable[..., User], product_factory: Callable[..., Product], - audit_log_factory, - audit_log_manager, + audit_log_factory: Callable[..., AuditLog], + audit_log_manager: AuditLogManager, ): p1 = product_factory() p2 = product_factory() @@ -205,18 +217,29 @@ class TestAuditLog: class TestIPRecordManager: - def test_init(self, thl_web_rr: PostgresConfig, thl_redis_config: RedisConfig, ip_record_manager): - instance = IPRecordManager(pg_config=thl_web_rr: PostgresConfig, redis_config=thl_redis_config) + def test_init( + self, + thl_web_rr: PostgresConfig, + thl_redis_config: RedisConfig, + ip_record_manager: IPRecordManager, + ): + instance = IPRecordManager(pg_config=thl_web_rr, redis_config=thl_redis_config) assert isinstance(instance, IPRecordManager) assert isinstance(ip_record_manager, IPRecordManager) - def test_create(self, ip_record_manager, user, ip_information): + def test_create( + self, + ip_record_manager: IPRecordManager, + user: User, + ip_information: IPInformation, + ): instance = ip_record_manager.create_dummy( user_id=user.user_id, ip=ip_information.ip ) assert isinstance(instance, IPRecord) assert isinstance(instance.forwarded_ips, list) + assert isinstance(instance.forwarded_ip_records, list) assert isinstance(instance.forwarded_ip_records[0], IPRecord) assert isinstance(instance.forwarded_ips[0], str) @@ -228,10 +251,10 @@ class TestIPRecordManager: def test_prefetch_info( self, - ip_record_factory, - ip_information_factory, - ip_geoname, - user, + ip_record_factory: Callable[..., IPRecord], + ip_information_factory: Callable[..., IPInformation], + ip_geoname: IPGeoname, + user: User, thl_web_rr: PostgresConfig, thl_redis_config: RedisConfig, ): @@ -239,15 +262,17 @@ class TestIPRecordManager: ip = fake.ipv4_public() ip_information_factory(ip=ip, geoname=ip_geoname) ipr: IPRecord = ip_record_factory(user_id=user.user_id, ip=ip) + assert isinstance(ipr, IPRecord) assert ipr.information is None assert len(ipr.forwarded_ip_records) >= 1 + assert isinstance(ipr.forwarded_ip_records, list) fipr = ipr.forwarded_ip_records[0] assert fipr.information is None ipr.prefetch_ipinfo( - pg_config=thl_web_rr: PostgresConfig, - redis_config=thl_redis_config: RedisConfig, + pg_config=thl_web_rr, + redis_config=thl_redis_config, include_forwarded=True, ) assert isinstance(ipr.information, GeoIPInformation) @@ -256,8 +281,8 @@ class TestIPRecordManager: ip_information_factory(ip=fipr.ip, geoname=ip_geoname) ipr.prefetch_ipinfo( - pg_config=thl_web_rr: PostgresConfig, - redis_config=thl_redis_config: RedisConfig, + pg_config=thl_web_rr, + redis_config=thl_redis_config, include_forwarded=True, ) assert fipr.information is not None @@ -265,28 +290,35 @@ class TestIPRecordManager: @pytest.mark.usefixtures("user_iphistory_manager_clear_cache") class TestUserIpHistoryManager: - def test_init(self, thl_web_rr: PostgresConfig, thl_redis_config: RedisConfig, user_iphistory_manager): + def test_init( + self, + thl_web_rr: PostgresConfig, + thl_redis_config: RedisConfig, + user_iphistory_manager: UserIpHistoryManager, + ): instance = UserIpHistoryManager( - pg_config=thl_web_rr: PostgresConfig, redis_config=thl_redis_config + pg_config=thl_web_rr, redis_config=thl_redis_config ) assert isinstance(instance, UserIpHistoryManager) assert isinstance(user_iphistory_manager, UserIpHistoryManager) def test_latest_record( self, - user_iphistory_manager, - user, - ip_record_factory, - ip_information_factory, - ip_geoname, + user_iphistory_manager: UserIpHistoryManager, + user: User, + ip_record_factory: Callable[..., IPRecord], + ip_information_factory: Callable[..., IPInformation], + ip_geoname: IPGeoname, ): ip = fake.ipv4_public() ip_information_factory(ip=ip, geoname=ip_geoname, is_anonymous=True) ipr1: IPRecord = ip_record_factory(user_id=user.user_id, ip=ip) ipr = user_iphistory_manager.get_user_latest_ip_record(user=user) + assert isinstance(ipr, IPRecord) assert ipr.ip == ipr1.ip assert ipr.is_anonymous + assert isinstance(ipr.information, GeoIPInformation) assert ipr.information.lookup_prefix == "/32" ip = fake.ipv6() @@ -294,7 +326,9 @@ class TestUserIpHistoryManager: ipr2: IPRecord = ip_record_factory(user_id=user.user_id, ip=ip) ipr = user_iphistory_manager.get_user_latest_ip_record(user=user) + assert isinstance(ipr, IPRecord) assert ipr.ip == ipr2.ip + assert isinstance(ipr.information, GeoIPInformation) assert ipr.information.lookup_prefix == "/64" assert ipr.information is not None assert not ipr.is_anonymous @@ -303,6 +337,8 @@ class TestUserIpHistoryManager: assert country_iso == ip_geoname.country_iso iph = user_iphistory_manager.get_user_ip_history(user_id=user.user_id) + assert isinstance(iph, UserIPHistory) + assert isinstance(iph.ips, list) assert iph.ips[0].information is not None assert iph.ips[1].information is not None assert iph.ips[0].country_iso == country_iso @@ -310,7 +346,12 @@ class TestUserIpHistoryManager: assert iph.ips[0].ip == ipr1.ip assert iph.ips[1].ip == ipr2.ip - def test_virgin(self, user, user_iphistory_manager, ip_record_factory): + def test_virgin( + self, + user: User, + user_iphistory_manager: UserIpHistoryManager, + ip_record_factory: Callable[..., IPRecord], + ): iph = user_iphistory_manager.get_user_ip_history(user_id=user.user_id) assert len(iph.ips) == 0 @@ -320,16 +361,18 @@ class TestUserIpHistoryManager: def test_out_of_order( self, - ip_record_factory, - user, - user_iphistory_manager, - ip_information_factory, - ip_geoname, + ip_record_factory: Callable[..., IPRecord], + user: User, + user_iphistory_manager: UserIpHistoryManager, + ip_information_factory: Callable[..., IPInformation], + ip_geoname: IPGeoname, ): # Create the user-ip association BEFORE the ip even exists in the ipinfo table ip = fake.ipv4_public() ip_record_factory(user_id=user.user_id, ip=ip) iph = user_iphistory_manager.get_user_ip_history(user_id=user.user_id) + assert isinstance(iph, UserIPHistory) + assert isinstance(iph.ips, list) assert len(iph.ips) == 1 ipr = iph.ips[0] assert ipr.information is None @@ -337,6 +380,8 @@ class TestUserIpHistoryManager: ip_information_factory(ip=ip, geoname=ip_geoname, is_anonymous=True) iph = user_iphistory_manager.get_user_ip_history(user_id=user.user_id) + assert isinstance(iph, UserIPHistory) + assert isinstance(iph.ips, list) assert len(iph.ips) == 1 ipr = iph.ips[0] assert ipr.information is not None @@ -344,16 +389,18 @@ class TestUserIpHistoryManager: def test_out_of_order_ipv6( self, - ip_record_factory, - user, - user_iphistory_manager, - ip_information_factory, - ip_geoname, + ip_record_factory: Callable[..., IPRecord], + user: User, + user_iphistory_manager: UserIpHistoryManager, + ip_information_factory: Callable[..., IPInformation], + ip_geoname: IPGeoname, ): # Create the user-ip association BEFORE the ip even exists in the ipinfo table ip = fake.ipv6() ip_record_factory(user_id=user.user_id, ip=ip) iph = user_iphistory_manager.get_user_ip_history(user_id=user.user_id) + assert isinstance(iph, UserIPHistory) + assert isinstance(iph.ips, list) assert len(iph.ips) == 1 ipr = iph.ips[0] assert ipr.information is None @@ -361,6 +408,8 @@ class TestUserIpHistoryManager: ip_information_factory(ip=ip, geoname=ip_geoname, is_anonymous=True) iph = user_iphistory_manager.get_user_ip_history(user_id=user.user_id) + assert isinstance(iph, UserIPHistory) + assert isinstance(iph.ips, list) assert len(iph.ips) == 1 ipr = iph.ips[0] assert ipr.information is not None diff --git a/tests/managers/thl/test_wall_manager.py b/tests/managers/thl/test_wall_manager.py index 1071627..b8a636f 100644 --- a/tests/managers/thl/test_wall_manager.py +++ b/tests/managers/thl/test_wall_manager.py @@ -1,24 +1,36 @@ +from __future__ import annotations + +from collections.abc import Callable from datetime import UTC, datetime, timedelta from decimal import Decimal from uuid import uuid4 import pytest +from pydantic import PositiveInt +from generalresearch.managers.thl.session import SessionManager +from generalresearch.managers.thl.wall import WallCacheManager, WallManager from generalresearch.models import Source from generalresearch.models.thl.session import ( ReportValue, + Session, Status, StatusCode1, ) +from generalresearch.models.thl.user import User class TestWallManager: @pytest.mark.parametrize("wall_count", [1, 2, 5, 10, 50, 99]) def test_get_wall_events( - self, wall_manager, session_factory, user, wall_count, utc_hour_ago + self, + wall_manager: WallManager, + session_factory: Callable[..., Session], + user: User, + wall_count: PositiveInt, + utc_hour_ago: datetime, ): - from generalresearch.models.thl.session import Session s1: Session = session_factory( user=user, wall_count=wall_count, started=utc_hour_ago @@ -62,12 +74,15 @@ class TestWallManager: ] def test_get_wall_events_list_input( - self, wall_manager, session_factory, user, utc_hour_ago + self, + wall_manager: WallManager, + session_factory: Callable[..., Session], + user: User, + utc_hour_ago: datetime, ): - from generalresearch.models.thl.session import Session session_ids = [] - for idx in range(10): + for _ in range(10): s: Session = session_factory(user=user, wall_count=5, started=utc_hour_ago) session_ids.append(s.id) @@ -82,7 +97,7 @@ class TestWallManager: assert session_ids == res1 - def test_create_wall(self, wall_manager, user, session): + def test_create_wall(self, wall_manager: WallManager, user: User, session: Session): w = wall_manager.create( session_id=session.id, user_id=user.user_id, @@ -98,7 +113,13 @@ class TestWallManager: w2 = wall_manager.get_from_uuid(wall_uuid=w.uuid) assert w == w2 - def test_report_wall_abandon(self, wall_manager, user, session, utc_hour_ago): + def test_report_wall_abandon( + self, + wall_manager: WallManager, + user: User, + session: Session, + utc_hour_ago: datetime, + ): w1 = wall_manager.create( session_id=session.id, user_id=user.user_id, @@ -138,7 +159,12 @@ class TestWallManager: # the status and finished get updated def test_report_wall( - self, wall_manager, session_manager, user, session, utc_hour_ago + self, + wall_manager: WallManager, + session_manager: SessionManager, + user: User, + session: Session, + utc_hour_ago: datetime, ): w1 = wall_manager.create( session_id=session.id, @@ -174,7 +200,13 @@ class TestWallManager: assert Status.COMPLETE == w2.status assert "This survey blows!" == w2.report_notes - def test_filter_wall_attempts(self, wall_manager, user, session, utc_hour_ago): + def test_filter_wall_attempts( + self, + wall_manager: WallManager, + user: User, + session: Session, + utc_hour_ago: datetime, + ): res = wall_manager.filter_wall_attempts(user_id=user.user_id) assert len(res) == 0 wall_manager.create( @@ -205,12 +237,16 @@ class TestWallManager: class TestWallCacheManager: - def test_get_attempts_none(self, wall_cache_manager, user): + def test_get_attempts_none(self, wall_cache_manager: WallCacheManager, user: User): attempts = wall_cache_manager.get_attempts(user.user_id) assert len(attempts) == 0 def test_get_wall_events( - self, wall_cache_manager, wall_manager, session_manager, user + self, + wall_cache_manager: WallCacheManager, + wall_manager: WallManager, + session_manager: SessionManager, + user: User, ): start1 = datetime.now(UTC) - timedelta(hours=3) start2 = datetime.now(UTC) - timedelta(hours=2) diff --git a/tests/models/custom_types/test_dsn.py b/tests/models/custom_types/test_dsn.py index eb27526..d8c7c53 100644 --- a/tests/models/custom_types/test_dsn.py +++ b/tests/models/custom_types/test_dsn.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from uuid import uuid4 import pytest diff --git a/tests/models/custom_types/test_therest.py b/tests/models/custom_types/test_therest.py index 13e9bae..01bc644 100644 --- a/tests/models/custom_types/test_therest.py +++ b/tests/models/custom_types/test_therest.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import json from uuid import UUID diff --git a/tests/models/dynata/test_eligbility.py b/tests/models/dynata/test_eligbility.py index 27de5b3..b3a9f13 100644 --- a/tests/models/dynata/test_eligbility.py +++ b/tests/models/dynata/test_eligbility.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from datetime import UTC, datetime diff --git a/tests/models/gr/test_authentication.py b/tests/models/gr/test_authentication.py index d4db112..d2a7054 100644 --- a/tests/models/gr/test_authentication.py +++ b/tests/models/gr/test_authentication.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import binascii import json import os @@ -7,9 +9,14 @@ from random import randint from uuid import uuid4 import pytest +from redis import Redis -from generalresearch.models.gr.authentication import GRUser +from generalresearch.models.gr.authentication import Claims, GRToken, GRUser +from generalresearch.models.gr.business import Business from generalresearch.models.gr.team import Membership, Team +from generalresearch.models.thl.product import Product +from generalresearch.pg_helper import PostgresConfig +from generalresearch.redis_helper import RedisConfig SSO_ISSUER = "" @@ -29,7 +36,13 @@ class TestGRUser: def test_businesses(self): pass - def test_teams(self, gr_user: GRUser, membership, gr_db, gr_redis_config): + def test_teams( + self, + gr_user: GRUser, + membership: Membership, + gr_db: PostgresConfig, + gr_redis_config: RedisConfig, + ): assert gr_user.teams is None @@ -41,15 +54,15 @@ class TestGRUser: def test_prefetch_team_duplicates( self, - gr_user_token, + gr_user_token: GRToken, gr_user: GRUser, membership: Membership, product_factory: Callable[..., Product], - membership_factory, + membership_factory: Callable[..., Membership], team: Team, thl_web_rr: PostgresConfig, - gr_redis_config, - gr_db, + gr_redis_config: RedisConfig, + gr_db: PostgresConfig, ): product_factory(team=team) membership_factory(team=team, gr_user=gr_user) @@ -67,9 +80,9 @@ class TestGRUser: product_factory: Callable[..., Product], team: Team, membership: Membership, - gr_db, + gr_db: PostgresConfig, thl_web_rr: PostgresConfig, - gr_redis_config, + gr_redis_config: RedisConfig, ): from generalresearch.models.thl.product import Product @@ -78,6 +91,8 @@ class TestGRUser: # Create a new Team membership, and then create a Product that # is part of that team membership.prefetch_team(pg_config=gr_db, redis_config=gr_redis_config) + assert isinstance(membership.team, Team) + p: Product = product_factory(team=team) assert p.id_int assert team.uuid == membership.team.uuid @@ -87,7 +102,7 @@ class TestGRUser: gr_user.prefetch_products( pg_config=gr_db, - thl_pg_config=thl_web_rr: PostgresConfig, + thl_pg_config=thl_web_rr, redis_config=gr_redis_config, ) assert isinstance(gr_user.products, list) @@ -97,7 +112,7 @@ class TestGRUser: class TestGRUserMethods: - def test_cache_key(self, gr_user, gr_redis): + def test_cache_key(self, gr_user: GRUser, gr_redis: RedisConfig): assert isinstance(gr_user.cache_key, str) assert ":" in gr_user.cache_key assert str(gr_user.id) in gr_user.cache_key @@ -105,11 +120,11 @@ class TestGRUserMethods: def test_to_redis( self, gr_user: GRUser, - gr_redis, + gr_redis: Redis, team: Team, business: Business, product_factory: Callable[..., Product], - membership_factory: Callable[Membership], + membership_factory: Callable[..., Membership], ): product_factory(team=team, business=business) membership_factory(team=team, gr_user=gr_user) @@ -125,11 +140,11 @@ class TestGRUserMethods: def test_set_cache( self, gr_user: GRUser, - gr_user_token, - gr_redis, - gr_db, + gr_user_token: GRToken, + gr_redis: Redis, + gr_db: PostgresConfig, thl_web_rr: PostgresConfig, - gr_redis_config, + gr_redis_config: RedisConfig, ): assert gr_redis.get(name=gr_user.cache_key) is None assert gr_redis.get(name=f"{gr_user.cache_key}:team_uuids") is None @@ -137,7 +152,7 @@ class TestGRUserMethods: assert gr_redis.get(name=f"{gr_user.cache_key}:product_uuids") is None gr_user.set_cache( - pg_config=gr_db, thl_web_rr=thl_web_rr: PostgresConfig, redis_config=gr_redis_config + pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config ) assert gr_redis.get(name=gr_user.cache_key) is not None @@ -148,14 +163,14 @@ class TestGRUserMethods: def test_set_cache_gr_user( self, gr_user: GRUser, - gr_user_token, - gr_redis, - gr_redis_config, - gr_db, + gr_user_token: GRToken, + gr_redis: RedisConfig, + gr_redis_config: RedisConfig, + gr_db: PostgresConfig, thl_web_rr: PostgresConfig, product_factory: Callable[..., Product], - team, - membership_factory, + team: Team, + membership_factory: Callable[..., Membership], thl_redis_config: RedisConfig, ): from generalresearch.models.gr.authentication import GRUser @@ -164,7 +179,7 @@ class TestGRUserMethods: membership_factory(team=team, gr_user=gr_user) gr_user.set_cache( - pg_config=gr_db, thl_web_rr=thl_web_rr: PostgresConfig, redis_config=gr_redis_config + pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config ) res: str = gr_redis.get(name=gr_user.cache_key) @@ -176,27 +191,27 @@ class TestGRUserMethods: gru2.prefetch_products( pg_config=gr_db, - thl_pg_config=thl_web_rr: PostgresConfig, - redis_config=thl_redis_config: RedisConfig, + thl_pg_config=thl_web_rr, + redis_config=thl_redis_config, ) assert gru2.product_uuids == [p1.uuid] def test_set_cache_team_uuids( self, - gr_user, - membership, - gr_user_token, - gr_redis, - gr_db, + gr_user: GRUser, + membership: Membership, + gr_user_token: GRToken, + gr_redis: Redis, + gr_db: PostgresConfig, thl_web_rr: PostgresConfig, product_factory: Callable[..., Product], - team, - gr_redis_config, + team: Team, + gr_redis_config: RedisConfig, ): product_factory(team=team) gr_user.set_cache( - pg_config=gr_db, thl_web_rr=thl_web_rr: PostgresConfig, redis_config=gr_redis_config + pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config ) res = json.loads(gr_redis.get(name=f"{gr_user.cache_key}:team_uuids")) assert len(res) == 1 @@ -206,18 +221,18 @@ class TestGRUserMethods: def test_set_cache_business_uuids( self, gr_user: GRUser, - gr_redis, - gr_db, + gr_redis: Redis, + gr_db: PostgresConfig, thl_web_rr: PostgresConfig, product_factory: Callable[..., Product], business: Business, - team, - gr_redis_config, + team: Team, + gr_redis_config: RedisConfig, ): product_factory(team=team, business=business) gr_user.set_cache( - pg_config=gr_db, thl_web_rr=thl_web_rr: PostgresConfig, redis_config=gr_redis_config + pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config ) res = json.loads(gr_redis.get(name=f"{gr_user.cache_key}:business_uuids")) assert len(res) == 1 @@ -225,20 +240,20 @@ class TestGRUserMethods: def test_set_cache_product_uuids( self, - gr_user, - membership, - gr_user_token, - gr_redis, - gr_db, + gr_user: GRUser, + membership: Membership, + gr_user_token: GRToken, + gr_redis: Redis, + gr_db: PostgresConfig, thl_web_rr: PostgresConfig, product_factory: Callable[..., Product], - team, - gr_redis_config, + team: Team, + gr_redis_config: RedisConfig, ): product_factory(team=team) gr_user.set_cache( - pg_config=gr_db, thl_web_rr=thl_web_rr: PostgresConfig, redis_config=gr_redis_config + pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config ) res = json.loads(gr_redis.get(name=f"{gr_user.cache_key}:product_uuids")) assert len(res) == 1 @@ -248,9 +263,7 @@ class TestGRUserMethods: class TestGRToken: @pytest.fixture - def gr_token(self, gr_user): - from generalresearch.models.gr.authentication import GRToken - + def gr_token(self, gr_user: GRUser): now = datetime.now(tz=UTC) token = binascii.hexlify(os.urandom(20)).decode() @@ -258,29 +271,26 @@ class TestGRToken: return gr_token - def test_init(self, gr_token): - from generalresearch.models.gr.authentication import GRToken - + def test_init(self, gr_token: GRToken): assert isinstance(gr_token, GRToken) assert gr_token.created - def test_user(self, gr_token, gr_db, gr_redis_config): - from generalresearch.models.gr.authentication import GRUser - + def test_user( + self, gr_token: GRToken, gr_db: PostgresConfig, gr_redis_config: RedisConfig + ): assert gr_token.user is None gr_token.prefetch_user(pg_config=gr_db, redis_config=gr_redis_config) assert isinstance(gr_token.user, GRUser) - def test_auth_header(self, gr_token): + def test_auth_header(self, gr_token: GRToken): assert isinstance(gr_token.auth_header, dict) class TestClaims: def test_init(self): - from generalresearch.models.gr.authentication import Claims d = { "iss": SSO_ISSUER, diff --git a/tests/models/gr/test_base.py b/tests/models/gr/test_base.py index 8da28d3..412fa52 100644 --- a/tests/models/gr/test_base.py +++ b/tests/models/gr/test_base.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import subprocess from collections.abc import Callable from pathlib import Path diff --git a/tests/models/gr/test_business.py b/tests/models/gr/test_business.py index 5239ac2..f0107de 100644 --- a/tests/models/gr/test_business.py +++ b/tests/models/gr/test_business.py @@ -19,6 +19,10 @@ from pytest import approx from generalresearch.currency import USDCent from generalresearch.incite.base import GRLDatasets +from generalresearch.incite.collections.thl_web import ( + SessionDFCollection, + WallDFCollection, +) from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge from generalresearch.managers.gr.business import BusinessBankAccountManager from generalresearch.managers.gr.team import TeamManager @@ -29,7 +33,7 @@ from generalresearch.managers.thl.payout import ( PayoutEventManager, ) from generalresearch.models.gr.business import ( - business: Business, + Business, BusinessAddress, BusinessBankAccount, BusinessContact, @@ -50,7 +54,7 @@ class TestBusinessBankAccount: def test_init( self, - business: business: Business, + business: Business, business_bank_account_manager: BusinessBankAccountManager, ): from generalresearch.models.gr.business import ( @@ -68,7 +72,7 @@ class TestBusinessBankAccount: def test_business( self, business_bank_account: BusinessBankAccount, - business: business: Business, + business: Business, gr_db: PostgresConfig, gr_redis_config: RedisConfig, ): @@ -79,7 +83,7 @@ class TestBusinessBankAccount: business_bank_account.prefetch_business( pg_config=gr_db, redis_config=gr_redis_config ) - assert isinstance(business_bank_account.business: Business, Business) + assert isinstance(business_bank_account.business, Business) assert business_bank_account.business.uuid == business.uuid @@ -112,13 +116,13 @@ class TestBusiness: def test_init(self, business: Business): - assert isinstance(business: Business, Business) + assert isinstance(business, Business) assert isinstance(business.id, int) assert isinstance(business.uuid, str) def test_str_and_repr( self, - business: business: Business, + business: Business, product_factory: Callable[..., Product], thl_web_rr: PostgresConfig, ledger_manager: LedgerManager, @@ -181,12 +185,12 @@ class TestBusiness: pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) business.prebuild_payouts( - thl_pg_config=thl_web_rr: PostgresConfig, + thl_pg_config=thl_web_rr, thl_lm=thl_ledger_manager, bpem=business_payout_event_manager, ) business.prebuild_balance( - thl_pg_config=thl_web_rr: PostgresConfig, + thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -198,7 +202,7 @@ class TestBusiness: def test_addresses( self, - business: business: Business, + business: Business, business_address: BusinessAddress, gr_db: PostgresConfig, ): @@ -213,7 +217,7 @@ class TestBusiness: def test_teams( self, - business: business: Business, + business: Business, team: Team, team_manager: TeamManager, gr_db: PostgresConfig, @@ -231,7 +235,7 @@ class TestBusiness: def test_products( self, - business: business: Business, + business: Business, product_factory: Callable[..., Product], thl_web_rr: PostgresConfig, ): @@ -254,7 +258,7 @@ class TestBusiness: business.prefetch_products(thl_pg_config=thl_web_rr) assert len(business.products) == 3 - def test_bank_accounts(self, business: business: Business, gr_db: PostgresConfig): + def test_bank_accounts(self, business: Business, gr_db: PostgresConfig): assert business.products is None # It's an empty list after prefetch @@ -264,7 +268,7 @@ class TestBusiness: def test_balance( self, - business: business: Business, + business: Business, mnt_filepath: GRLDatasets, client_no_amm: DaskClient, thl_web_rr: PostgresConfig, @@ -275,7 +279,7 @@ class TestBusiness: with pytest.raises(expected_exception=AssertionError) as cm: business.prebuild_balance( - thl_pg_config=thl_web_rr: PostgresConfig, + thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -289,7 +293,7 @@ class TestBusiness: def test_payouts_no_accounts( self, - business: business: Business, + business: Business, product_factory: Callable[..., Product], thl_web_rr: PostgresConfig, thl_ledger_manager: ThlLedgerManager, @@ -299,7 +303,7 @@ class TestBusiness: with pytest.raises(expected_exception=AssertionError) as cm: business.prebuild_payouts( - thl_pg_config=thl_web_rr: PostgresConfig, + thl_pg_config=thl_web_rr, thl_lm=thl_ledger_manager, bpem=business_payout_event_manager, ) @@ -309,7 +313,7 @@ class TestBusiness: thl_ledger_manager.get_account_or_create_bp_wallet(product=p) business.prebuild_payouts( - thl_pg_config=thl_web_rr: PostgresConfig, + thl_pg_config=thl_web_rr, thl_lm=thl_ledger_manager, bpem=business_payout_event_manager, ) @@ -318,7 +322,7 @@ class TestBusiness: def test_payouts( self, - business: business: Business, + business: Business, product_factory: Callable[..., Product], bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], thl_ledger_manager: ThlLedgerManager, @@ -338,7 +342,7 @@ class TestBusiness: ) business.prebuild_payouts( - thl_pg_config=thl_web_rr: PostgresConfig, + thl_pg_config=thl_web_rr, thl_lm=thl_ledger_manager, bpem=business_payout_event_manager, ) @@ -356,7 +360,7 @@ class TestBusiness: thl_lm=thl_ledger_manager ) business.prebuild_payouts( - thl_pg_config=thl_web_rr: PostgresConfig, + thl_pg_config=thl_web_rr, thl_lm=thl_ledger_manager, bpem=business_payout_event_manager, ) @@ -367,12 +371,12 @@ class TestBusiness: def test_payouts_totals( self, - business: business: Business, + business: Business, product_factory: Callable[..., Product], bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], thl_ledger_manager: ThlLedgerManager, thl_web_rr: PostgresConfig, - business_payout_event_manager, + business_payout_event_manager: BusinessPayoutEventManager, create_main_accounts: Callable[..., None], ): @@ -406,7 +410,7 @@ class TestBusiness: ) business.prebuild_payouts( - thl_pg_config=thl_web_rr: PostgresConfig, + thl_pg_config=thl_web_rr, thl_lm=thl_ledger_manager, bpem=business_payout_event_manager, ) @@ -419,7 +423,7 @@ class TestBusiness: def test_pop_financial( self, - business: business: Business, + business: Business, thl_web_rr: PostgresConfig, thl_ledger_manager: ThlLedgerManager, mnt_filepath: GRLDatasets, @@ -428,7 +432,7 @@ class TestBusiness: ): assert business.pop_financial is None business.prebuild_pop_financial( - thl_pg_config=thl_web_rr: PostgresConfig, + thl_pg_config=thl_web_rr, thl_lm=thl_ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -438,7 +442,7 @@ class TestBusiness: def test_bp_accounts( self, - business: business: Business, + business: Business, thl_web_rr: PostgresConfig, product_factory: Callable[..., Product], thl_ledger_manager: ThlLedgerManager, @@ -480,7 +484,7 @@ class TestBusinessBalance: def test_single_product( self, - business: business: Business, + business: Business, product_factory: Callable[..., Product], user_factory: Callable[..., User], mnt_filepath, @@ -519,7 +523,7 @@ class TestBusinessBalance: pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) business.prebuild_balance( - thl_pg_config=thl_web_rr: PostgresConfig, + thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -541,7 +545,7 @@ class TestBusinessBalance: def test_multi_product( self, - business: business: Business, + business: Business, product_factory: Callable[..., Product], user_factory: Callable[..., User], mnt_filepath: GRLDatasets, @@ -579,7 +583,7 @@ class TestBusinessBalance: pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) business.prebuild_balance( - thl_pg_config=thl_web_rr: PostgresConfig, + thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -625,7 +629,7 @@ class TestBusinessBalance: def test_multi_product_multi_payout( self, - business: business: Business, + business: Business, product_factory: Callable[..., Product], user_factory: Callable[..., User], mnt_filepath: GRLDatasets, @@ -665,7 +669,7 @@ class TestBusinessBalance: payout_event_manager.set_account_lookup_table(thl_lm=thl_ledger_manager) bp_payout_factory( - product=u1.product: Product, + product=u1.product, amount=USDCent(5), created=start + timedelta(days=4), skip_wallet_balance_check=True, @@ -673,7 +677,7 @@ class TestBusinessBalance: ) bp_payout_factory( - product=u2.product: Product, + product=u2.product, amount=USDCent(50), created=start + timedelta(days=4), skip_wallet_balance_check=True, @@ -684,7 +688,7 @@ class TestBusinessBalance: pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) business.prebuild_balance( - thl_pg_config=thl_web_rr: PostgresConfig, + thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -699,7 +703,7 @@ class TestBusinessBalance: def test_multi_product_multi_payout_adjustment( self, - business: business: Business, + business: Business, product_factory: Callable[..., Product], user_factory: Callable[..., User], mnt_filepath: GRLDatasets, @@ -758,7 +762,7 @@ class TestBusinessBalance: payout_event_manager.set_account_lookup_table(thl_lm=thl_ledger_manager) bp_payout_factory( - product=u1.product: Product, + product=u1.product, amount=USDCent(250), created=start + timedelta(days=3), skip_wallet_balance_check=True, @@ -766,7 +770,7 @@ class TestBusinessBalance: ) bp_payout_factory( - product=u2.product: Product, + product=u2.product, amount=USDCent(50), created=start + timedelta(days=4), skip_wallet_balance_check=True, @@ -796,7 +800,7 @@ class TestBusinessBalance: assert df.shape == (20, 28) business.prebuild_balance( - thl_pg_config=thl_web_rr: PostgresConfig, + thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -833,7 +837,7 @@ class TestBusinessBalance: create_main_accounts: Callable[..., None], delete_df_collection: Callable[..., None], ledger_collection, - business: business: Business, + business: Business, user_factory: Callable[..., User], product_factory: Callable[..., Product], session_with_tx_factory: Callable[..., Session], @@ -869,7 +873,7 @@ class TestBusinessBalance: ) payout_event_manager.set_account_lookup_table(thl_lm=thl_ledger_manager) bp_payout_factory( - product=u1.product: Product, + product=u1.product, amount=USDCent(71), ext_ref_id=uuid4().hex, created=start + timedelta(days=1, minutes=1), @@ -898,7 +902,7 @@ class TestBusinessBalance: pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) business.prebuild_balance( - thl_pg_config=thl_web_rr: PostgresConfig, + thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -946,7 +950,7 @@ class TestBusinessBalance: def test_multi_product_multi_payout_adjustment_at_timestamp( self, - business: business: Business, + business: Business, product_factory: Callable[..., Product], user_factory: Callable[..., User], mnt_filepath: GRLDatasets, @@ -1022,7 +1026,7 @@ class TestBusinessBalance: payout_event_manager.set_account_lookup_table(thl_lm=thl_ledger_manager) bp_payout_factory( - product=u1.product: Product, + product=u1.product, amount=USDCent(250), created=start + timedelta(days=3), skip_wallet_balance_check=True, @@ -1030,7 +1034,7 @@ class TestBusinessBalance: ) bp_payout_factory( - product=u2.product: Product, + product=u2.product, amount=USDCent(50), created=start + timedelta(days=4), skip_wallet_balance_check=True, @@ -1060,7 +1064,7 @@ class TestBusinessBalance: assert df.shape == (20, 28) business.prebuild_balance( - thl_pg_config=thl_web_rr: PostgresConfig, + thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -1068,7 +1072,7 @@ class TestBusinessBalance: ) business.prebuild_balance( - thl_pg_config=thl_web_rr: PostgresConfig, + thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -1078,7 +1082,7 @@ class TestBusinessBalance: day1_bal = business.balance business.prebuild_balance( - thl_pg_config=thl_web_rr: PostgresConfig, + thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -1088,7 +1092,7 @@ class TestBusinessBalance: day2_bal = business.balance business.prebuild_balance( - thl_pg_config=thl_web_rr: PostgresConfig, + thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -1098,7 +1102,7 @@ class TestBusinessBalance: day3_bal = business.balance business.prebuild_balance( - thl_pg_config=thl_web_rr: PostgresConfig, + thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -1108,7 +1112,7 @@ class TestBusinessBalance: day4_bal = business.balance business.prebuild_balance( - thl_pg_config=thl_web_rr: PostgresConfig, + thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -1118,7 +1122,7 @@ class TestBusinessBalance: day5_bal = business.balance business.prebuild_balance( - thl_pg_config=thl_web_rr: PostgresConfig, + thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -1183,7 +1187,7 @@ class TestBusinessMethods: def test_set_cache( self, - business: business: Business, + business: Business, gr_redis: RedisConfig, gr_db: PostgresConfig, thl_web_rr: PostgresConfig, @@ -1219,7 +1223,7 @@ class TestBusinessMethods: business.set_cache( pg_config=gr_db, - thl_web_rr=thl_web_rr: PostgresConfig, + thl_web_rr=thl_web_rr, redis_config=gr_redis_config, client=client_no_amm, ds=mnt_filepath, @@ -1245,7 +1249,7 @@ class TestBusinessMethods: def test_set_cache_business( self, - business: business: Business, + business: Business, gr_db: PostgresConfig, thl_web_rr: PostgresConfig, product_factory: Callable[..., Product], @@ -1282,7 +1286,7 @@ class TestBusinessMethods: business.set_cache( pg_config=gr_db, - thl_web_rr=thl_web_rr: PostgresConfig, + thl_web_rr=thl_web_rr, redis_config=gr_redis_config, client=client_no_amm, ds=mnt_filepath, @@ -1345,15 +1349,15 @@ class TestBusinessMethods: self, enriched_session_merge, client_no_amm: DaskClient, - wall_collection, - session_collection, + wall_collection: WallDFCollection, + session_collection: SessionDFCollection, thl_web_rr: PostgresConfig, user_factory: Callable[..., User], start: datetime, session_factory: Callable[..., Session], product_factory: Callable[..., Product], delete_df_collection: Callable[..., None], - business: business: Business, + business: Business, mnt_filepath: GRLDatasets, mnt_gr_api_dir: Path, ): @@ -1380,11 +1384,11 @@ class TestBusinessMethods: client=client_no_amm, session_coll=session_collection, wall_coll=wall_collection, - pg_config=thl_web_rr: PostgresConfig, + pg_config=thl_web_rr, ) business.prebuild_enriched_session_parquet( - thl_pg_config=thl_web_rr: PostgresConfig, + thl_pg_config=thl_web_rr, ds=mnt_filepath, client=client_no_amm, mnt_gr_api=mnt_gr_api_dir, @@ -1401,15 +1405,15 @@ class TestBusinessMethods: self, enriched_wall_merge, client_no_amm: DaskClient, - wall_collection, - session_collection, + wall_collection: WallDFCollection, + session_collection: SessionDFCollection, thl_web_rr: PostgresConfig, user_factory: Callable[..., User], start: datetime, session_factory: Callable[..., Session], product_factory: Callable[..., Product], delete_df_collection: Callable[..., None], - business: business: Business, + business: Business, mnt_filepath: GRLDatasets, mnt_gr_api_dir: Path, ): @@ -1436,11 +1440,11 @@ class TestBusinessMethods: client=client_no_amm, session_coll=session_collection, wall_coll=wall_collection, - pg_config=thl_web_rr: PostgresConfig, + pg_config=thl_web_rr, ) business.prebuild_enriched_wall_parquet( - thl_pg_config=thl_web_rr: PostgresConfig, + thl_pg_config=thl_web_rr, ds=mnt_filepath, client=client_no_amm, mnt_gr_api=mnt_gr_api_dir, diff --git a/tests/models/gr/test_team.py b/tests/models/gr/test_team.py index 26300b9..dc7d4b9 100644 --- a/tests/models/gr/test_team.py +++ b/tests/models/gr/test_team.py @@ -97,7 +97,7 @@ class TestTeam: def test_businesses( self, team: Team, - business: business: Business, + business: Business, team_manager: TeamManager, gr_db: PostgresConfig, gr_redis_config: RedisConfig, @@ -160,7 +160,7 @@ class TestTeamMethods: team.set_cache( pg_config=gr_db, - thl_web_rr=thl_web_rr: PostgresConfig, + thl_web_rr=thl_web_rr, redis_config=gr_redis_config, client=client_no_amm, ds=mnt_filepath, @@ -192,7 +192,7 @@ class TestTeamMethods: team.set_cache( pg_config=gr_db, - thl_web_rr=thl_web_rr: PostgresConfig, + thl_web_rr=thl_web_rr, redis_config=gr_redis_config, client=client_no_amm, ds=mnt_filepath, @@ -254,11 +254,11 @@ class TestTeamMethods: client=client_no_amm, session_coll=session_collection, wall_coll=wall_collection, - pg_config=thl_web_rr: PostgresConfig, + pg_config=thl_web_rr, ) team.prebuild_enriched_session_parquet( - thl_pg_config=thl_web_rr: PostgresConfig, + thl_pg_config=thl_web_rr, ds=mnt_filepath, client=client_no_amm, mnt_gr_api=mnt_gr_api_dir, @@ -310,11 +310,11 @@ class TestTeamMethods: client=client_no_amm, session_coll=session_collection, wall_coll=wall_collection, - pg_config=thl_web_rr: PostgresConfig, + pg_config=thl_web_rr, ) team.prebuild_enriched_wall_parquet( - thl_pg_config=thl_web_rr: PostgresConfig, + thl_pg_config=thl_web_rr, ds=mnt_filepath, client=client_no_amm, mnt_gr_api=mnt_gr_api_dir, -- cgit v1.2.3 From cf239865ce440e1a71ee2360514eaeb018620ac9 Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Wed, 26 Aug 2026 16:46:45 -0700 Subject: Ruff afternoon! --- generalresearch/models/network/nmap/execute.py | 2 +- generalresearch/models/thl/soft_pair.py | 3 +- generalresearch/models/thl/survey/condition.py | 11 +- test_utils/managers/conftest.py | 2 +- test_utils/models/thl/conftest.py | 162 +++++++++++++++++++++ test_utils/precision/__init__.py | 0 test_utils/precision/conftest.py | 129 ++++++++++++++++ test_utils/spectrum/conftest.py | 58 +++++++- tests/conftest.py | 3 + tests/models/innovate/test_question.py | 2 + .../models/legacy/test_offerwall_parse_response.py | 2 + tests/models/legacy/test_profiling_questions.py | 6 +- .../models/legacy/test_user_question_answer_in.py | 54 +++---- tests/models/morning/test.py | 2 + tests/models/network/test_mtr.py | 5 +- tests/models/network/test_nmap.py | 11 +- tests/models/network/test_nmap_parser.py | 12 +- tests/models/network/test_rdns.py | 2 + tests/models/precision/__init__.py | 115 --------------- tests/models/precision/test_survey.py | 42 +++--- tests/models/prodege/test_survey_participation.py | 17 +-- tests/models/spectrum/test_question.py | 5 + tests/models/spectrum/test_survey.py | 48 +++--- tests/models/spectrum/test_survey_manager.py | 103 +++++-------- tests/models/test_currency.py | 2 + tests/models/test_device.py | 8 +- tests/models/test_finance.py | 30 ++-- tests/models/thl/question/test_question_info.py | 137 +---------------- tests/models/thl/question/test_user_info.py | 29 +--- tests/models/thl/test_adjustments.py | 58 +++++--- tests/models/thl/test_bucket.py | 8 +- tests/models/thl/test_buyer.py | 2 + tests/models/thl/test_contest/test_contest.py | 2 + .../thl/test_contest/test_leaderboard_contest.py | 27 +++- .../models/thl/test_contest/test_raffle_contest.py | 39 ++++- tests/models/thl/test_ledger.py | 2 + tests/models/thl/test_marketplace_condition.py | 38 +---- tests/models/thl/test_payout.py | 4 +- tests/models/thl/test_payout_format.py | 2 + tests/models/thl/test_product.py | 86 ++++++----- tests/models/thl/test_product_userwalletconfig.py | 8 +- tests/models/thl/test_soft_pair.py | 10 +- tests/models/thl/test_upkquestion.py | 78 ++++------ tests/models/thl/test_user.py | 76 +++------- tests/models/thl/test_user_iphistory.py | 2 + tests/models/thl/test_user_metadata.py | 2 + tests/models/thl/test_user_streak.py | 2 + tests/models/thl/test_wall.py | 2 + tests/models/thl/test_wall_session.py | 2 + tests/wall_status_codes/test_analyze.py | 2 + tests/wxet/models/test_definitions.py | 37 ++--- tests/wxet/models/test_finish_type.py | 2 + 52 files changed, 785 insertions(+), 708 deletions(-) create mode 100644 test_utils/precision/__init__.py create mode 100644 test_utils/precision/conftest.py (limited to 'test_utils/models') diff --git a/generalresearch/models/network/nmap/execute.py b/generalresearch/models/network/nmap/execute.py index 8a73307..6c05c89 100644 --- a/generalresearch/models/network/nmap/execute.py +++ b/generalresearch/models/network/nmap/execute.py @@ -19,7 +19,7 @@ def execute_nmap( enable_advanced: bool = True, timing: int = 4, scan_group_id: UUIDStr | None = None, -): +) -> NmapRun: config = NmapRunCommand( options=NmapRunCommandOptions( top_ports=top_ports, diff --git a/generalresearch/models/thl/soft_pair.py b/generalresearch/models/thl/soft_pair.py index 6ff1165..f3b2b6f 100644 --- a/generalresearch/models/thl/soft_pair.py +++ b/generalresearch/models/thl/soft_pair.py @@ -4,6 +4,7 @@ from dataclasses import dataclass from enum import Enum from generalresearch.models import Source +from generalresearch.models.dynata.survey import DynataCondition from generalresearch.models.thl.survey.condition import ( MarketplaceCondition, ) @@ -34,7 +35,7 @@ class SoftPairResult: pair_type: SoftPairResultType source: Source survey_id: str - conditions: set[MarketplaceCondition] | None = None + conditions: set[MarketplaceCondition | DynataCondition] | None = None @property def survey_sid(self) -> str: diff --git a/generalresearch/models/thl/survey/condition.py b/generalresearch/models/thl/survey/condition.py index a85073c..514ee64 100644 --- a/generalresearch/models/thl/survey/condition.py +++ b/generalresearch/models/thl/survey/condition.py @@ -248,7 +248,7 @@ class MarketplaceCondition(BaseModel, ABC): return d @staticmethod - def is_numeric_including_inf(s) -> bool: + def is_numeric_including_inf(s: Any) -> bool: try: float(s) return True @@ -263,10 +263,9 @@ class MarketplaceCondition(BaseModel, ABC): # Fancy repr that only shows the first and last 3 values if there are more than 6. repr_args = list(self.__repr_args__()) for n, (k, v) in enumerate(repr_args): - if k == "values": - if v and len(v) > 6: - v = v[:3] + ["…"] + v[-3:] - repr_args[n] = ("values", v) + if k == "values" and v and len(v) > 6: + v = v[:3] + ["…"] + v[-3:] + repr_args[n] = ("values", v) join_str = ", " repr_str = join_str.join( repr(v) if a is None else f"{a}={v!r}" for a, v in repr_args @@ -334,5 +333,5 @@ class MarketplaceCondition(BaseModel, ABC): except ValueError: return None values = self.values_ranges - passes = any([start <= x <= end for start, end in values for x in answer]) + passes = any(start <= x <= end for start, end in values for x in answer) return not passes if self.negate else passes diff --git a/test_utils/managers/conftest.py b/test_utils/managers/conftest.py index e5d6015..b03a646 100644 --- a/test_utils/managers/conftest.py +++ b/test_utils/managers/conftest.py @@ -195,7 +195,7 @@ def setup_cashoutmethod_db( @pytest.fixture(scope="session") -def spectrum_manager(spectrum_rw: SqlHelper) -> SpectrumSurveyManager: +def spectrum_survey_manager(spectrum_rw: SqlHelper) -> SpectrumSurveyManager: from generalresearch.managers.spectrum.survey import ( SpectrumSurveyManager, ) diff --git a/test_utils/models/thl/conftest.py b/test_utils/models/thl/conftest.py index 5b21c6b..907306d 100644 --- a/test_utils/models/thl/conftest.py +++ b/test_utils/models/thl/conftest.py @@ -433,3 +433,165 @@ def auditlog_factory(audit_log_manager: AuditLogManager): ) return _inner + + +@pytest.fixture(scope="session") +def profiling_info_json() -> str: + return ( + '[{"property_label": "hispanic", "cardinality": "*", "prop_type": "i", "country_iso": "us", ' + '"property_id": "05170ae296ab49178a075cab2a2073a6", "item_id": "7911ec1468b146ee870951f8ae9cbac1", ' + '"item_label": "panamanian", "gold_standard": 1, "options": [{"id": "c358c11e72c74fa2880358f1d4be85ab", ' + '"label": "not_hispanic"}, {"id": "b1d6c475770849bc8e0200054975dc9c", "label": "yes_hispanic"}, ' + '{"id": "bd1eb44495d84b029e107c188003c2bd", "label": "other_hispanic"}, ' + '{"id": "f290ad5e75bf4f4ea94dc847f57c1bd3", "label": "mexican"}, ' + '{"id": "49f50f2801bd415ea353063bfc02d252", "label": "puerto_rican"}, ' + '{"id": "dcbe005e522f4b10928773926601f8bf", "label": "cuban"}, ' + '{"id": "467ef8ddb7ac4edb88ba9ef817cbb7e9", "label": "salvadoran"}, ' + '{"id": "3c98e7250707403cba2f4dc7b877c963", "label": "dominican"}, ' + '{"id": "981ee77f6d6742609825ef54fea824a8", "label": "guatemalan"}, ' + '{"id": "81c8057b809245a7ae1b8a867ea6c91e", "label": "colombian"}, ' + '{"id": "513656d5f9e249fa955c3b527d483b93", "label": "honduran"}, ' + '{"id": "afc8cddd0c7b4581bea24ccd64db3446", "label": "ecuadorian"}, ' + '{"id": "61f34b36e80747a89d85e1eb17536f84", "label": "argentinian"}, ' + '{"id": "5330cfa681d44aa8ade3a6d0ea198e44", "label": "peruvian"}, ' + '{"id": "e7bceaffd76e486596205d8545019448", "label": "nicaraguan"}, ' + '{"id": "b7bbb2ebf8424714962e6c4f43275985", "label": "spanish"}, ' + '{"id": "8bf539785e7a487892a2f97e52b1932d", "label": "venezuelan"}, ' + '{"id": "7911ec1468b146ee870951f8ae9cbac1", "label": "panamanian"}], "category": [{"id": ' + '"4fd8381d5a1c4409ab007ca254ced084", "label": "Demographic", "path": "/Demographic", ' + '"adwords_vertical_id": null}]}, {"property_label": "ethnic_group", "cardinality": "*", "prop_type": ' + '"i", "country_iso": "us", "property_id": "15070958225d4132b7f6674fcfc979f6", "item_id": ' + '"64b7114cf08143949e3bcc3d00a5d8a0", "item_label": "other_ethnicity", "gold_standard": 1, "options": [{' + '"id": "a72e97f4055e4014a22bee4632cbf573", "label": "caucasians"}, ' + '{"id": "4760353bc0654e46a928ba697b102735", "label": "black_or_african_american"}, ' + '{"id": "20ff0a2969fa4656bbda5c3e0874e63b", "label": "asian"}, ' + '{"id": "107e0a79e6b94b74926c44e70faf3793", "label": "native_hawaiian_or_other_pacific_islander"}, ' + '{"id": "900fa12691d5458c8665bf468f1c98c1", "label": "native_americans"}, ' + '{"id": "64b7114cf08143949e3bcc3d00a5d8a0", "label": "other_ethnicity"}], "category": [{"id": ' + '"4fd8381d5a1c4409ab007ca254ced084", "label": "Demographic", "path": "/Demographic", ' + '"adwords_vertical_id": null}]}, {"property_label": "educational_attainment", "cardinality": "?", ' + '"prop_type": "i", "country_iso": "us", "property_id": "2637783d4b2b4075b93e2a156e16e1d8", "item_id": ' + '"934e7b81d6744a1baa31bbc51f0965d5", "item_label": "other_education", "gold_standard": 1, "options": [{' + '"id": "df35ef9e474b4bf9af520aa86630202d", "label": "3rd_grade_completion"}, ' + '{"id": "83763370a1064bd5ba76d1b68c4b8a23", "label": "8th_grade_completion"}, ' + '{"id": "f0c25a0670c340bc9250099dcce50957", "label": "not_high_school_graduate"}, ' + '{"id": "02ff74c872bd458983a83847e1a9f8fd", "label": "high_school_completion"}, ' + '{"id": "ba8beb807d56441f8fea9b490ed7561c", "label": "vocational_program_completion"}, ' + '{"id": "65373a5f348a410c923e079ddbb58e9b", "label": "some_college_completion"}, ' + '{"id": "2d15d96df85d4cc7b6f58911fdc8d5e2", "label": "associate_academic_degree_completion"}, ' + '{"id": "497b1fedec464151b063cd5367643ffa", "label": "bachelors_degree_completion"}, ' + '{"id": "295133068ac84424ae75e973dc9f2a78", "label": "some_graduate_completion"}, ' + '{"id": "e64f874faeff4062a5aa72ac483b4b9f", "label": "masters_degree_completion"}, ' + '{"id": "cbaec19a636d476385fb8e7842b044f5", "label": "doctorate_degree_completion"}, ' + '{"id": "934e7b81d6744a1baa31bbc51f0965d5", "label": "other_education"}], "category": [{"id": ' + '"4fd8381d5a1c4409ab007ca254ced084", "label": "Demographic", "path": "/Demographic", ' + '"adwords_vertical_id": null}]}, {"property_label": "household_spoken_language", "cardinality": "*", ' + '"prop_type": "i", "country_iso": "us", "property_id": "5a844571073d482a96853a0594859a51", "item_id": ' + '"62b39c1de141422896ad4ab3c4318209", "item_label": "dut", "gold_standard": 1, "options": [{"id": ' + '"f65cd57b79d14f0f8460761ce41ec173", "label": "ara"}, {"id": "6d49de1f8f394216821310abd29392d9", ' + '"label": "zho"}, {"id": "be6dc23c2bf34c3f81e96ddace22800d", "label": "eng"}, ' + '{"id": "ddc81f28752d47a3b1c1f3b8b01a9b07", "label": "fre"}, {"id": "2dbb67b29bd34e0eb630b1b8385542ca", ' + '"label": "ger"}, {"id": "a747f96952fc4b9d97edeeee5120091b", "label": "hat"}, ' + '{"id": "7144b04a3219433baac86273677551fa", "label": "hin"}, {"id": "e07ff3e82c7149eaab7ea2b39ee6a6dc", ' + '"label": "ita"}, {"id": "b681eff81975432ebfb9f5cc22dedaa3", "label": "jpn"}, ' + '{"id": "5cb20440a8f64c9ca62fb49c1e80cdef", "label": "kor"}, {"id": "171c4b77d4204bc6ac0c2b81e38a10ff", ' + '"label": "pan"}, {"id": "8c3ec18e6b6c4a55a00dd6052e8e84fb", "label": "pol"}, ' + '{"id": "3ce074d81d384dd5b96f1fb48f87bf01", "label": "por"}, {"id": "6138dc951990458fa88a666f6ddd907b", ' + '"label": "rus"}, {"id": "e66e5ecc07df4ebaa546e0b436f034bd", "label": "spa"}, ' + '{"id": "5a981b3d2f0d402a96dd2d0392ec2fcb", "label": "tgl"}, {"id": "b446251bd211403487806c4d0a904981", ' + '"label": "vie"}, {"id": "92fb3ee337374e2db875fb23f52eed46", "label": "xxx"}, ' + '{"id": "8b1f590f12f24cc1924d7bdcbe82081e", "label": "ind"}, {"id": "bf3f4be556a34ff4b836420149fd2037", ' + '"label": "tur"}, {"id": "87ca815c43ba4e7f98cbca98821aa508", "label": "zul"}, ' + '{"id": "0adbf915a7a64d67a87bb3ce5d39ca54", "label": "may"}, {"id": "62b39c1de141422896ad4ab3c4318209", ' + '"label": "dut"}], "category": [{"id": "4fd8381d5a1c4409ab007ca254ced084", "label": "Demographic", ' + '"path": "/Demographic", "adwords_vertical_id": null}]}, {"property_label": "gender", "cardinality": ' + '"?", "prop_type": "i", "country_iso": "us", "property_id": "73175402104741549f21de2071556cd7", ' + '"item_id": "093593e316344cd3a0ac73669fca8048", "item_label": "other_gender", "gold_standard": 1, ' + '"options": [{"id": "b9fc5ea07f3a4252a792fd4a49e7b52b", "label": "male"}, ' + '{"id": "9fdb8e5e18474a0b84a0262c21e17b56", "label": "female"}, ' + '{"id": "093593e316344cd3a0ac73669fca8048", "label": "other_gender"}], "category": [{"id": ' + '"4fd8381d5a1c4409ab007ca254ced084", "label": "Demographic", "path": "/Demographic", ' + '"adwords_vertical_id": null}]}, {"property_label": "age_in_years", "cardinality": "?", "prop_type": ' + '"n", "country_iso": "us", "property_id": "94f7379437874076b345d76642d4ce6d", "item_id": null, ' + '"item_label": null, "gold_standard": 1, "category": [{"id": "4fd8381d5a1c4409ab007ca254ced084", ' + '"label": "Demographic", "path": "/Demographic", "adwords_vertical_id": null}]}, {"property_label": ' + '"children_age_gender", "cardinality": "*", "prop_type": "i", "country_iso": "us", "property_id": ' + '"e926142fcea94b9cbbe13dc7891e1e7f", "item_id": "b7b8074e95334b008e8958ccb0a204f1", "item_label": ' + '"female_18", "gold_standard": 1, "options": [{"id": "16a6448ec24c48d4993d78ebee33f9b4", ' + '"label": "male_under_1"}, {"id": "809c04cb2e3b4a3bbd8077ab62cdc220", "label": "female_under_1"}, ' + '{"id": "295e05bb6a0843bc998890b24c99841e", "label": "no_children"}, ' + '{"id": "142cb948d98c4ae8b0ef2ef10978e023", "label": "male_0"}, ' + '{"id": "5a5c1b0e9abc48a98b3bc5f817d6e9d0", "label": "male_1"}, ' + '{"id": "286b1a9afb884bdfb676dbb855479d1e", "label": "male_2"}, ' + '{"id": "942ca3cda699453093df8cbabb890607", "label": "male_3"}, ' + '{"id": "995818d432f643ec8dd17e0809b24b56", "label": "male_4"}, ' + '{"id": "f38f8b57f25f4cdea0f270297a1e7a5c", "label": "male_5"}, ' + '{"id": "975df709e6d140d1a470db35023c432d", "label": "male_6"}, ' + '{"id": "f60bd89bbe0f4e92b90bccbc500467c2", "label": "male_7"}, ' + '{"id": "6714ceb3ed5042c0b605f00b06814207", "label": "male_8"}, ' + '{"id": "c03c2f8271d443cf9df380e84b4dea4c", "label": "male_9"}, ' + '{"id": "11690ee0f5a54cb794f7ddd010d74fa2", "label": "male_10"}, ' + '{"id": "17bef9a9d14b4197b2c5609fa94b0642", "label": "male_11"}, ' + '{"id": "e79c8338fe28454f89ccc78daf6f409a", "label": "male_12"}, ' + '{"id": "3a4f87acb3fa41f4ae08dfe2858238c1", "label": "male_13"}, ' + '{"id": "36ffb79d8b7840a7a8cb8d63bbc8df59", "label": "male_14"}, ' + '{"id": "1401a508f9664347aee927f6ec5b0a40", "label": "male_15"}, ' + '{"id": "6e0943c5ec4a4f75869eb195e3eafa50", "label": "male_16"}, ' + '{"id": "47d4b27b7b5242758a9fff13d3d324cf", "label": "male_17"}, ' + '{"id": "9ce886459dd44c9395eb77e1386ab181", "label": "female_0"}, ' + '{"id": "6499ccbf990d4be5b686aec1c7353fd8", "label": "female_1"}, ' + '{"id": "d85ceaa39f6d492abfc8da49acfd14f2", "label": "female_2"}, ' + '{"id": "18edb45c138e451d8cb428aefbb80f9c", "label": "female_3"}, ' + '{"id": "bac6f006ed9f4ccf85f48e91e99fdfd1", "label": "female_4"}, ' + '{"id": "5a6a1a8ad00c4ce8be52dcb267b034ff", "label": "female_5"}, ' + '{"id": "6bff0acbf6364c94ad89507bcd5f4f45", "label": "female_6"}, ' + '{"id": "d0d56a0a6b6f4516a366a2ce139b4411", "label": "female_7"}, ' + '{"id": "bda6028468044b659843e2bef4db2175", "label": "female_8"}, ' + '{"id": "dbb6d50325464032b456357b1a6e5e9c", "label": "female_9"}, ' + '{"id": "b87a93d7dc1348edac5e771684d63fb8", "label": "female_10"}, ' + '{"id": "11449d0d98f14e27ba47de40b18921d7", "label": "female_11"}, ' + '{"id": "16156501e97b4263962cbbb743840292", "label": "female_12"}, ' + '{"id": "04ee971c89a345cc8141a45bce96050c", "label": "female_13"}, ' + '{"id": "e818d310bfbc4faba4355e5d2ed49d4f", "label": "female_14"}, ' + '{"id": "440d25e078924ba0973163153c417ed6", "label": "female_15"}, ' + '{"id": "78ff804cc9b441c5a524bd91e3d1f8bf", "label": "female_16"}, ' + '{"id": "4b04d804d7d84786b2b1c22e4ed440f5", "label": "female_17"}, ' + '{"id": "28bc848cd3ff44c3893c76bfc9bc0c4e", "label": "male_18"}, ' + '{"id": "b7b8074e95334b008e8958ccb0a204f1", "label": "female_18"}], "category": [{"id": ' + '"e18ba6e9d51e482cbb19acf2e6f505ce", "label": "Parenting", "path": "/People & Society/Family & ' + 'Relationships/Family/Parenting", "adwords_vertical_id": "58"}]}, {"property_label": "home_postal_code", ' + '"cardinality": "?", "prop_type": "x", "country_iso": "us", "property_id": ' + '"f3b32ebe78014fbeb1ed6ff77d6338bf", "item_id": null, "item_label": null, "gold_standard": 1, ' + '"category": [{"id": "4fd8381d5a1c4409ab007ca254ced084", "label": "Demographic", "path": "/Demographic", ' + '"adwords_vertical_id": null}]}, {"property_label": "household_income", "cardinality": "?", "prop_type": ' + '"n", "country_iso": "us", "property_id": "ff5b1d4501d5478f98de8c90ef996ac1", "item_id": null, ' + '"item_label": null, "gold_standard": 1, "category": [{"id": "4fd8381d5a1c4409ab007ca254ced084", ' + '"label": "Demographic", "path": "/Demographic", "adwords_vertical_id": null}]}]' + ) + + +@pytest.fixture(scope="session") +def profiling_user_info_json() -> str: + return ( + '{"user_profile_knowledge": [], "marketplace_profile_knowledge": [{"source": "d", "question_id": ' + '"1", "answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "pr", ' + '"question_id": "3", "answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": ' + '"h", "question_id": "60", "answer": ["58"], "created": "2023-11-07T16:41:05.234096Z"}, ' + '{"source": "c", "question_id": "43", "answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, ' + '{"source": "s", "question_id": "211", "answer": ["111"], "created": ' + '"2023-11-07T16:41:05.234096Z"}, {"source": "s", "question_id": "1843", "answer": ["111"], ' + '"created": "2023-11-07T16:41:05.234096Z"}, {"source": "h", "question_id": "13959", "answer": [' + '"244155"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "c", "question_id": "33092", ' + '"answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "c", "question_id": "gender", ' + '"answer": ["10682"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "e", "question_id": ' + '"gender", "answer": ["male"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "f", ' + '"question_id": "gender", "answer": ["male"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": ' + '"i", "question_id": "gender", "answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, ' + '{"source": "c", "question_id": "137510", "answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, ' + '{"source": "m", "question_id": "gender", "answer": ["1"], "created": ' + '"2023-11-07T16:41:05.234096Z"}, {"source": "o", "question_id": "gender", "answer": ["male"], ' + '"created": "2023-11-07T16:41:05.234096Z"}, {"source": "c", "question_id": "gender_plus", "answer": [' + '"7657644"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "i", "question_id": ' + '"gender_plus", "answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "c", ' + '"question_id": "income_level", "answer": ["9071"], "created": "2023-11-07T16:41:05.234096Z"}]}' + ) diff --git a/test_utils/precision/__init__.py b/test_utils/precision/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/test_utils/precision/conftest.py b/test_utils/precision/conftest.py new file mode 100644 index 0000000..7acfe6a --- /dev/null +++ b/test_utils/precision/conftest.py @@ -0,0 +1,129 @@ +from typing import Any + +import pytest + + +@pytest.fixture(scope="session") +def precision_survey_json() -> dict[str, Any]: + return { + "cpi": "1.44", + "country_isos": "ca", + "language_isos": "eng", + "country_iso": "ca", + "language_iso": "eng", + "buyer_id": "7047", + "bid_loi": 1200, + "bid_ir": 0.45, + "source": "e", + "used_question_ids": ["age", "country_iso", "gender", "gender_1"], + "survey_id": "0000", + "group_id": "633473", + "status": "open", + "name": "beauty survey", + "survey_guid": "c7f375c5077d4c6c8209ff0b539d7183", + "category_id": "-1", + "global_conversion": None, + "desired_count": 96, + "achieved_count": 0, + "allowed_devices": "1,2,3", + "entry_link": "https://www.opinionetwork.com/survey/entry.aspx?mid=[%MID%]&project=633473&key=%%key%%", + "excluded_surveys": "470358,633286", + "quotas": [ + { + "name": "25-34,Male,Quebec", + "id": "2324110", + "guid": "23b5760d24994bc08de451b3e62e77c7", + "status": "open", + "desired_count": 48, + "achieved_count": 0, + "termination_count": 0, + "overquota_count": 0, + "condition_hashes": ["b41e1a3", "bc89ee8", "4124366", "9f32c61"], + }, + { + "name": "25-34,Female,Quebec", + "id": "2324111", + "guid": "0706f1a88d7e4f11ad847c03012e68d2", + "status": "open", + "desired_count": 48, + "achieved_count": 0, + "termination_count": 4, + "overquota_count": 0, + "condition_hashes": ["b41e1a3", "0cdc304", "500af2c", "9f32c61"], + }, + ], + "conditions": { + "b41e1a3": { + "logical_operator": "OR", + "value_type": 1, + "negate": False, + "question_id": "country_iso", + "values": ["ca"], + "criterion_hash": "b41e1a3", + "value_len": 1, + "sizeof": 2, + }, + "bc89ee8": { + "logical_operator": "OR", + "value_type": 1, + "negate": False, + "question_id": "gender", + "values": ["male"], + "criterion_hash": "bc89ee8", + "value_len": 1, + "sizeof": 4, + }, + "4124366": { + "logical_operator": "OR", + "value_type": 1, + "negate": False, + "question_id": "gender_1", + "values": ["male"], + "criterion_hash": "4124366", + "value_len": 1, + "sizeof": 4, + }, + "9f32c61": { + "logical_operator": "OR", + "value_type": 1, + "negate": False, + "question_id": "age", + "values": ["25", "26", "27", "28", "29", "30", "31", "32", "33", "34"], + "criterion_hash": "9f32c61", + "value_len": 10, + "sizeof": 20, + }, + "0cdc304": { + "logical_operator": "OR", + "value_type": 1, + "negate": False, + "question_id": "gender", + "values": ["female"], + "criterion_hash": "0cdc304", + "value_len": 1, + "sizeof": 6, + }, + "500af2c": { + "logical_operator": "OR", + "value_type": 1, + "negate": False, + "question_id": "gender_1", + "values": ["female"], + "criterion_hash": "500af2c", + "value_len": 1, + "sizeof": 6, + }, + }, + "expected_end_date": "2024-06-28T10:40:33.000000Z", + "created": None, + "updated": None, + "is_live": True, + "all_hashes": [ + "0cdc304", + "b41e1a3", + "9f32c61", + "bc89ee8", + "4124366", + "500af2c", + ], + } diff --git a/test_utils/spectrum/conftest.py b/test_utils/spectrum/conftest.py index eb2e289..d186a5b 100644 --- a/test_utils/spectrum/conftest.py +++ b/test_utils/spectrum/conftest.py @@ -1,7 +1,10 @@ +from __future__ import annotations + import logging import time from datetime import UTC, datetime -from typing import TYPE_CHECKING +from decimal import Decimal +from typing import TYPE_CHECKING, Any import pytest @@ -83,3 +86,56 @@ def setup_spectrum_surveys( conn.commit() # Wait a second to make sure the spectrum-grpc pulls these from the db into global-vars time.sleep(1) + + +@pytest.fixture(scope="session") +def spectrum_api_survey_json() -> dict[str, Any]: + return { + "survey_id": 29333264, + "survey_name": "#29333264", + "survey_status": 22, + "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=UTC), + "mod_on": datetime(2024, 5, 20, 18, 18, 31, tzinfo=UTC), + "soft_launch": False, + "click_balancing": 0, + "price_type": 1, + "pii": False, + "buyer_message": "", + "buyer_id": 4726, + "incl_excl": 0, + "cpi": Decimal("1.20"), + "last_complete_date": None, + "project_last_complete_date": None, + "quotas": [ + { + "quota_id": "c2bc961e-4f26-4223-b409-ebe9165cfdf5", + "quantities": {"currently_open": 491, "remaining": 495, "achieved": 0}, + "criteria": [ + { + "qualification_code": 214, + "range_sets": [{"units": 311, "to": 64, "from": 18}], + } + ], + } + ], + "qualifications": [ + { + "range_sets": [{"units": 311, "to": 64, "from": 18}], + "qualification_code": 212, + }, + {"condition_codes": ["111", "117", "112"], "qualification_code": 1202}, + ], + "country_iso": "fr", + "language_iso": "fre", + "bid_ir": 0.4, + "bid_loi": 600, + "overall_ir": None, + "overall_loi": None, + "last_block_ir": None, + "last_block_loi": None, + "survey_exclusions": set(), + "exclusion_period": 0, + } diff --git a/tests/conftest.py b/tests/conftest.py index 6748592..4777e15 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -23,4 +23,7 @@ pytest_plugins = [ "test_utils.models.network.conftest", "test_utils.models.thl.conftest", "test_utils.models.upk.conftest", + # -- Marketplaces + "test_utils.precision.conftest", + "test_utils.spectrum.conftest", ] diff --git a/tests/models/innovate/test_question.py b/tests/models/innovate/test_question.py index b0c2964..b206177 100644 --- a/tests/models/innovate/test_question.py +++ b/tests/models/innovate/test_question.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from generalresearch.models import Source from generalresearch.models.innovate.question import ( InnovateQuestion, diff --git a/tests/models/legacy/test_offerwall_parse_response.py b/tests/models/legacy/test_offerwall_parse_response.py index b1c96ad..56ba077 100644 --- a/tests/models/legacy/test_offerwall_parse_response.py +++ b/tests/models/legacy/test_offerwall_parse_response.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import json from generalresearch.models import Source diff --git a/tests/models/legacy/test_profiling_questions.py b/tests/models/legacy/test_profiling_questions.py index 1afaa6b..6f781ae 100644 --- a/tests/models/legacy/test_profiling_questions.py +++ b/tests/models/legacy/test_profiling_questions.py @@ -1,7 +1,11 @@ +from __future__ import annotations + +from generalresearch.models.legacy.questions import UpkQuestionResponse + + class TestUpkQuestionResponse: def test_init(self): - from generalresearch.models.legacy.questions import UpkQuestionResponse s = ( '{"status": "success", "count": 7, "questions": [{"selector": "SL", "validation": {"patterns": [{' diff --git a/tests/models/legacy/test_user_question_answer_in.py b/tests/models/legacy/test_user_question_answer_in.py index 313862c..3fdaa05 100644 --- a/tests/models/legacy/test_user_question_answer_in.py +++ b/tests/models/legacy/test_user_question_answer_in.py @@ -1,9 +1,22 @@ +from __future__ import annotations + import json +from collections.abc import Callable +from datetime import datetime from decimal import Decimal from uuid import uuid4 import pytest +from generalresearch.managers.thl.user_manager.user_manager import UserManager +from generalresearch.models 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 + class TestUserQuestionAnswers: """This is for the GRS POST submission that may contain multiple @@ -15,21 +28,11 @@ class TestUserQuestionAnswers: def test_json_init( self, - product_manager: ProductManager, - user_manager, - session_manager, - wall_manager, user_factory: Callable[..., User], product: Product, - session_factory, - utc_hour_ago, + session_factory: Callable[..., Session], + utc_hour_ago: datetime, ): - from generalresearch.models import Source - from generalresearch.models.legacy.questions import ( - UserQuestionAnswers, - ) - from generalresearch.models.thl.session import Session, Wall - from generalresearch.models.thl.user import User u: User = user_factory(product=product) @@ -61,14 +64,7 @@ class TestUserQuestionAnswers: def test_simple_validation_errors( self, - product_manager: ProductManager, - user_manager, - session_manager, - wall_manager, ): - from generalresearch.models.legacy.questions import ( - UserQuestionAnswers, - ) with pytest.raises(ValueError): UserQuestionAnswers.model_validate( @@ -118,7 +114,7 @@ class TestUserQuestionAnswers: with pytest.raises(ValueError): answers = [ - {"question_id": uuid4().hex, "answer": ["a"]} for i in range(101) + {"question_id": uuid4().hex, "answer": ["a"]} for _ in range(101) ] UserQuestionAnswers.model_validate( { @@ -143,9 +139,6 @@ class TestUserQuestionAnswers: # TODO: depending on if or how many of these types of errors actually # occur, we could get fancy and just drop one of them. I don't # think this is worth exploring yet unless we see if it's a problem. - from generalresearch.models.legacy.questions import ( - UserQuestionAnswers, - ) consistent_qid = uuid4().hex with pytest.raises(ValueError) as cm: @@ -165,11 +158,11 @@ class TestUserQuestionAnswers: def test_allow_answer_failures_silent( self, - user_manager, + user_manager: UserManager, product: Product, user_factory: Callable[..., User], - utc_hour_ago, - session_factory, + utc_hour_ago: datetime, + session_factory: Callable[..., Session], ): """ There are many instances where suppliers may be submitting answers @@ -177,11 +170,6 @@ class TestUserQuestionAnswers: that one QuestionAnswerIn without "loosing" any of the other QuestionAnswerIn items that they provided. """ - from generalresearch.models.legacy.questions import ( - UserQuestionAnswers, - ) - from generalresearch.models.thl.session import Session, Wall - from generalresearch.models.thl.user import User u: User = user_factory(product=product) @@ -286,7 +274,7 @@ class TestUserQuestionAnswerIn: UserQuestionAnswerIn, ) - answer = [uuid4().hex[:6] for i in range(11)] + answer = [uuid4().hex[:6] for _ in range(11)] with pytest.raises(ValueError) as cm: UserQuestionAnswerIn.model_validate( {"question_id": uuid4().hex, "answer": answer} @@ -298,7 +286,7 @@ class TestUserQuestionAnswerIn: UserQuestionAnswerIn, ) - answer = ["aaa" for i in range(5)] + answer = ["aaa" for _ in range(5)] with pytest.raises(ValueError): UserQuestionAnswerIn.model_validate( {"question_id": uuid4().hex, "answer": answer} diff --git a/tests/models/morning/test.py b/tests/models/morning/test.py index 7474766..c1141fb 100644 --- a/tests/models/morning/test.py +++ b/tests/models/morning/test.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from datetime import UTC, datetime from generalresearch.models.morning.question import MorningQuestion diff --git a/tests/models/network/test_mtr.py b/tests/models/network/test_mtr.py index 840a773..7f8a736 100644 --- a/tests/models/network/test_mtr.py +++ b/tests/models/network/test_mtr.py @@ -1,12 +1,15 @@ +from __future__ import annotations + 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 fake = faker.Faker() -def test_execute_mtr(toolrun_manager): +def test_execute_mtr(toolrun_manager: ToolRunManager): ip = "65.19.129.53" run = execute_mtr(ip=ip, report_cycles=3) diff --git a/tests/models/network/test_nmap.py b/tests/models/network/test_nmap.py index a135a13..5e9f4d0 100644 --- a/tests/models/network/test_nmap.py +++ b/tests/models/network/test_nmap.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import subprocess import faker @@ -5,8 +7,8 @@ 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 PortState -from generalresearch.models.network.tool_run import ToolClass, ToolName +from generalresearch.models.network.nmap.result import NmapResult, PortState +from generalresearch.models.network.tool_run import NmapRun, Status, ToolClass, ToolName fake = faker.Faker() @@ -18,10 +20,13 @@ def resolve(host: str): def test_execute_nmap_scanme(toolrun_manager: ToolRunManager): ip = resolve("scanme.nmap.org") - run = execute_nmap(ip=ip, top_ports=None, ports="20-30", enable_advanced=False) + run: NmapRun = execute_nmap( + ip=ip, top_ports=None, ports="20-30", enable_advanced=False + ) assert run.tool_name == ToolName.NMAP assert run.tool_class == ToolClass.PORT_SCAN assert run.ip == ip + assert isinstance(run.parsed, NmapResult) result = run.parsed port22 = result._port_index[(IPProtocol.TCP, 22)] diff --git a/tests/models/network/test_nmap_parser.py b/tests/models/network/test_nmap_parser.py index 7822380..473a63f 100644 --- a/tests/models/network/test_nmap_parser.py +++ b/tests/models/network/test_nmap_parser.py @@ -1,8 +1,14 @@ +from __future__ import annotations + import os import pytest from generalresearch.models.network.nmap.parser import parse_nmap_xml +from generalresearch.models.network.nmap.result import ( + NmapResult, + NmapTrace, +) @pytest.fixture @@ -13,9 +19,11 @@ def nmap_raw_output_2(request) -> str: return data -def test_nmap_xml_parser(nmap_raw_output, nmap_raw_output_2): - n = parse_nmap_xml(nmap_raw_output) +def test_nmap_xml_parser(nmap_raw_output: str, nmap_raw_output_2: str): + n: NmapResult = parse_nmap_xml(nmap_raw_output) assert n.tcp_open_ports == [61232] + + assert isinstance(n.trace, NmapTrace) assert len(n.trace.hops) == 18 n = parse_nmap_xml(nmap_raw_output_2) diff --git a/tests/models/network/test_rdns.py b/tests/models/network/test_rdns.py index 5c3b024..1a15a28 100644 --- a/tests/models/network/test_rdns.py +++ b/tests/models/network/test_rdns.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import faker from generalresearch.managers.network.tool_run import ToolRunManager diff --git a/tests/models/precision/__init__.py b/tests/models/precision/__init__.py index 8006fa3..e69de29 100644 --- a/tests/models/precision/__init__.py +++ b/tests/models/precision/__init__.py @@ -1,115 +0,0 @@ -survey_json = { - "cpi": "1.44", - "country_isos": "ca", - "language_isos": "eng", - "country_iso": "ca", - "language_iso": "eng", - "buyer_id": "7047", - "bid_loi": 1200, - "bid_ir": 0.45, - "source": "e", - "used_question_ids": ["age", "country_iso", "gender", "gender_1"], - "survey_id": "0000", - "group_id": "633473", - "status": "open", - "name": "beauty survey", - "survey_guid": "c7f375c5077d4c6c8209ff0b539d7183", - "category_id": "-1", - "global_conversion": None, - "desired_count": 96, - "achieved_count": 0, - "allowed_devices": "1,2,3", - "entry_link": "https://www.opinionetwork.com/survey/entry.aspx?mid=[%MID%]&project=633473&key=%%key%%", - "excluded_surveys": "470358,633286", - "quotas": [ - { - "name": "25-34,Male,Quebec", - "id": "2324110", - "guid": "23b5760d24994bc08de451b3e62e77c7", - "status": "open", - "desired_count": 48, - "achieved_count": 0, - "termination_count": 0, - "overquota_count": 0, - "condition_hashes": ["b41e1a3", "bc89ee8", "4124366", "9f32c61"], - }, - { - "name": "25-34,Female,Quebec", - "id": "2324111", - "guid": "0706f1a88d7e4f11ad847c03012e68d2", - "status": "open", - "desired_count": 48, - "achieved_count": 0, - "termination_count": 4, - "overquota_count": 0, - "condition_hashes": ["b41e1a3", "0cdc304", "500af2c", "9f32c61"], - }, - ], - "conditions": { - "b41e1a3": { - "logical_operator": "OR", - "value_type": 1, - "negate": False, - "question_id": "country_iso", - "values": ["ca"], - "criterion_hash": "b41e1a3", - "value_len": 1, - "sizeof": 2, - }, - "bc89ee8": { - "logical_operator": "OR", - "value_type": 1, - "negate": False, - "question_id": "gender", - "values": ["male"], - "criterion_hash": "bc89ee8", - "value_len": 1, - "sizeof": 4, - }, - "4124366": { - "logical_operator": "OR", - "value_type": 1, - "negate": False, - "question_id": "gender_1", - "values": ["male"], - "criterion_hash": "4124366", - "value_len": 1, - "sizeof": 4, - }, - "9f32c61": { - "logical_operator": "OR", - "value_type": 1, - "negate": False, - "question_id": "age", - "values": ["25", "26", "27", "28", "29", "30", "31", "32", "33", "34"], - "criterion_hash": "9f32c61", - "value_len": 10, - "sizeof": 20, - }, - "0cdc304": { - "logical_operator": "OR", - "value_type": 1, - "negate": False, - "question_id": "gender", - "values": ["female"], - "criterion_hash": "0cdc304", - "value_len": 1, - "sizeof": 6, - }, - "500af2c": { - "logical_operator": "OR", - "value_type": 1, - "negate": False, - "question_id": "gender_1", - "values": ["female"], - "criterion_hash": "500af2c", - "value_len": 1, - "sizeof": 6, - }, - }, - "expected_end_date": "2024-06-28T10:40:33.000000Z", - "created": None, - "updated": None, - "is_live": True, - "all_hashes": ["0cdc304", "b41e1a3", "9f32c61", "bc89ee8", "4124366", "500af2c"], -} diff --git a/tests/models/precision/test_survey.py b/tests/models/precision/test_survey.py index ff2d6d1..4d671f2 100644 --- a/tests/models/precision/test_survey.py +++ b/tests/models/precision/test_survey.py @@ -1,10 +1,15 @@ -class TestPrecisionQuota: +from __future__ import annotations + +from typing import Any + +from generalresearch.models.precision import PrecisionStatus +from generalresearch.models.precision.survey import PrecisionSurvey - def test_quota_passes(self): - from generalresearch.models.precision.survey import PrecisionSurvey - from tests.models.precision import survey_json - s = PrecisionSurvey.model_validate(survey_json) +class TestPrecisionQuota: + + def test_quota_passes(self, precision_survey_json: dict[str, Any]): + s = PrecisionSurvey.model_validate(precision_survey_json) q = s.quotas[0] ce = {k: True for k in ["b41e1a3", "bc89ee8", "4124366", "9f32c61"]} assert q.matches(ce) @@ -16,12 +21,9 @@ class TestPrecisionQuota: assert not q.matches(ce) assert not q.matches({}) - def test_quota_passes_closed(self): - from generalresearch.models.precision import PrecisionStatus - from generalresearch.models.precision.survey import PrecisionSurvey - from tests.models.precision import survey_json + def test_quota_passes_closed(self, precision_survey_json: dict[str, Any]): - s = PrecisionSurvey.model_validate(survey_json) + s = PrecisionSurvey.model_validate(precision_survey_json) q = s.quotas[0] q.status = PrecisionStatus.CLOSED ce = {k: True for k in ["b41e1a3", "bc89ee8", "4124366", "9f32c61"]} @@ -32,20 +34,15 @@ class TestPrecisionQuota: class TestPrecisionSurvey: - def test_passes(self): - from generalresearch.models.precision.survey import PrecisionSurvey - from tests.models.precision import survey_json + def test_passes(self, precision_survey_json: dict[str, Any]): - s = PrecisionSurvey.model_validate(survey_json) + s = PrecisionSurvey.model_validate(precision_survey_json) ce = {k: True for k in ["b41e1a3", "bc89ee8", "4124366", "9f32c61"]} assert s.determine_eligibility(ce) - def test_elig_closed_quota(self): - from generalresearch.models.precision import PrecisionStatus - from generalresearch.models.precision.survey import PrecisionSurvey - from tests.models.precision import survey_json + def test_elig_closed_quota(self, precision_survey_json: dict[str, Any]): - s = PrecisionSurvey.model_validate(survey_json) + s = PrecisionSurvey.model_validate(precision_survey_json) ce = {k: True for k in ["b41e1a3", "bc89ee8", "4124366", "9f32c61"]} q = s.quotas[0] q.status = PrecisionStatus.CLOSED @@ -57,12 +54,9 @@ class TestPrecisionSurvey: # Now me match an open quota and dont match the closed quota, so we should be eligible assert s.determine_eligibility(ce) - def test_passes_sp(self): - from generalresearch.models.precision import PrecisionStatus - from generalresearch.models.precision.survey import PrecisionSurvey - from tests.models.precision import survey_json + def test_passes_sp(self, precision_survey_json: dict[str, Any]): - s = PrecisionSurvey.model_validate(survey_json) + s = PrecisionSurvey.model_validate(precision_survey_json) ce = {k: True for k in ["b41e1a3", "bc89ee8", "4124366", "9f32c61"]} passes, hashes = s.determine_eligibility_soft(ce) diff --git a/tests/models/prodege/test_survey_participation.py b/tests/models/prodege/test_survey_participation.py index e1ba9ab..10ce884 100644 --- a/tests/models/prodege/test_survey_participation.py +++ b/tests/models/prodege/test_survey_participation.py @@ -1,14 +1,17 @@ +from __future__ import annotations + from datetime import UTC, datetime, timedelta +from generalresearch.models.prodege import ProdegePastParticipationType +from generalresearch.models.prodege.survey import ( + ProdegePastParticipation, + ProdegeUserPastParticipation, +) + class TestProdegeParticipation: def test_exclude(self): - from generalresearch.models.prodege import ProdegePastParticipationType - from generalresearch.models.prodege.survey import ( - ProdegePastParticipation, - ProdegeUserPastParticipation, - ) now = datetime.now(tz=UTC) pp = ProdegePastParticipation.from_api( @@ -84,10 +87,6 @@ class TestProdegeParticipation: assert not pp.is_eligible(upps) def test_include(self): - from generalresearch.models.prodege.survey import ( - ProdegePastParticipation, - ProdegeUserPastParticipation, - ) now = datetime.now(tz=UTC) pp = ProdegePastParticipation.from_api( diff --git a/tests/models/spectrum/test_question.py b/tests/models/spectrum/test_question.py index 57d260d..a44286d 100644 --- a/tests/models/spectrum/test_question.py +++ b/tests/models/spectrum/test_question.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from datetime import UTC, datetime from generalresearch.models import Source @@ -32,6 +34,7 @@ class TestSpectrumQuestion: "mod_on": 1706557247467, } q = SpectrumQuestion.from_api(example_1, "us", "eng") + assert isinstance(q, SpectrumQuestion) expected_q = SpectrumQuestion( question_id="213", @@ -72,6 +75,8 @@ class TestSpectrumQuestion: "mod_on": 1706557249817, } q = SpectrumQuestion.from_api(example_2, "us", "eng") + assert isinstance(q, SpectrumQuestion) + expected_q = SpectrumQuestion( question_id="211", country_iso="us", diff --git a/tests/models/spectrum/test_survey.py b/tests/models/spectrum/test_survey.py index 7365c7e..7ddd407 100644 --- a/tests/models/spectrum/test_survey.py +++ b/tests/models/spectrum/test_survey.py @@ -1,15 +1,25 @@ +from __future__ import annotations + from datetime import UTC, datetime from decimal import Decimal +from generalresearch.models import ( + LogicalOperator, + Source, + TaskCalculationType, +) +from generalresearch.models.spectrum import SpectrumStatus +from generalresearch.models.spectrum.survey import ( + SpectrumCondition, + SpectrumQuota, + SpectrumSurvey, +) +from generalresearch.models.thl.survey.condition import ConditionValueType + class TestSpectrumCondition: def test_condition_create(self): - from generalresearch.models import LogicalOperator - from generalresearch.models.spectrum.survey import ( - SpectrumCondition, - ) - from generalresearch.models.thl.survey.condition import ConditionValueType c = SpectrumCondition.from_api( { @@ -64,10 +74,6 @@ class TestSpectrumCondition: class TestSpectrumQuota: def test_quota_create(self): - from generalresearch.models.spectrum.survey import ( - SpectrumCondition, - SpectrumQuota, - ) d = { "quota_id": "a846b545-4449-4d76-93a2-f8ebdf6e711e", @@ -84,9 +90,6 @@ class TestSpectrumQuota: assert q.is_open def test_quota_passes(self): - from generalresearch.models.spectrum.survey import ( - SpectrumQuota, - ) q = SpectrumQuota(remaining_count=57, condition_hashes=["a"]) assert q.passes({"a": True}) @@ -103,9 +106,6 @@ class TestSpectrumQuota: assert not q.passes({"a": True}) def test_quota_passes_soft(self): - from generalresearch.models.spectrum.survey import ( - SpectrumQuota, - ) q = SpectrumQuota(remaining_count=57, condition_hashes=["a", "b", "c"]) # Pass if we match all @@ -122,18 +122,6 @@ class TestSpectrumQuota: class TestSpectrumSurvey: def test_survey_create(self): - from generalresearch.models import ( - LogicalOperator, - Source, - TaskCalculationType, - ) - from generalresearch.models.spectrum import SpectrumStatus - from generalresearch.models.spectrum.survey import ( - SpectrumCondition, - SpectrumQuota, - SpectrumSurvey, - ) - from generalresearch.models.thl.survey.condition import ConditionValueType # Note: d is the raw response after calling SpectrumAPI.preprocess_survey() on it! d = { @@ -202,6 +190,8 @@ class TestSpectrumSurvey: "exclusion_period": 0, } s = SpectrumSurvey.from_api(d) + assert isinstance(s, SpectrumSurvey) + expected_survey = SpectrumSurvey( cpi=Decimal("1.20000"), country_isos=["fr"], @@ -303,6 +293,8 @@ class TestSpectrumSurvey: "exclusion_period": 0, } s = SpectrumSurvey.from_api(d) + assert isinstance(s, SpectrumSurvey) + assert {"212", "1202", "214"} == s.used_question_ids assert s.is_live assert s.is_open @@ -345,6 +337,8 @@ class TestSpectrumSurvey: "exclusion_period": 0, } s = SpectrumSurvey.from_api(d) + assert isinstance(s, SpectrumSurvey) + s.qualifications = ["a", "b", "c"] s.quotas = [ SpectrumQuota(remaining_count=10, condition_hashes=["a", "b"]), diff --git a/tests/models/spectrum/test_survey_manager.py b/tests/models/spectrum/test_survey_manager.py index ce26c44..11dc01f 100644 --- a/tests/models/spectrum/test_survey_manager.py +++ b/tests/models/spectrum/test_survey_manager.py @@ -1,69 +1,32 @@ -import copy +from __future__ import annotations + import logging from datetime import UTC, datetime from decimal import Decimal +from typing import Any from pymysql import IntegrityError -logger = logging.getLogger() +from generalresearch.config import is_debug +from generalresearch.managers.spectrum.survey import ( + SpectrumSurveyManager, +) +from generalresearch.sql_helper import SqlHelper -example_survey_api_response = { - "survey_id": 29333264, - "survey_name": "#29333264", - "survey_status": 22, - "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=UTC), - "mod_on": datetime(2024, 5, 20, 18, 18, 31, tzinfo=UTC), - "soft_launch": False, - "click_balancing": 0, - "price_type": 1, - "pii": False, - "buyer_message": "", - "buyer_id": 4726, - "incl_excl": 0, - "cpi": Decimal("1.20"), - "last_complete_date": None, - "project_last_complete_date": None, - "quotas": [ - { - "quota_id": "c2bc961e-4f26-4223-b409-ebe9165cfdf5", - "quantities": {"currently_open": 491, "remaining": 495, "achieved": 0}, - "criteria": [ - { - "qualification_code": 214, - "range_sets": [{"units": 311, "to": 64, "from": 18}], - } - ], - } - ], - "qualifications": [ - { - "range_sets": [{"units": 311, "to": 64, "from": 18}], - "qualification_code": 212, - }, - {"condition_codes": ["111", "117", "112"], "qualification_code": 1202}, - ], - "country_iso": "fr", - "language_iso": "fre", - "bid_ir": 0.4, - "bid_loi": 600, - "overall_ir": None, - "overall_loi": None, - "last_block_ir": None, - "last_block_loi": None, - "survey_exclusions": set(), - "exclusion_period": 0, -} +logger = logging.getLogger() class TestSpectrumSurvey: - def test_survey_create(self, settings, spectrum_manager, spectrum_rw): + def test_survey_create( + self, + spectrum_survey_manager: SpectrumSurveyManager, + spectrum_rw: SqlHelper, + spectrum_api_survey_json: dict[str, Any], + ): from generalresearch.models.spectrum.survey import SpectrumSurvey - assert settings.debug, "CRITICAL: Do not run this on production." + assert is_debug(), "CRITICAL: Do not run this on production." now = datetime.now(tz=UTC) spectrum_rw.execute_sql_query( @@ -73,24 +36,29 @@ class TestSpectrumSurvey: commit=True, ) - d = example_survey_api_response.copy() - s = SpectrumSurvey.from_api(d) - spectrum_manager.create(s) + s = SpectrumSurvey.from_api(spectrum_api_survey_json) + assert isinstance(s, SpectrumSurvey) + spectrum_survey_manager.create(s) - surveys = spectrum_manager.get_survey_library(updated_since=now) + surveys = spectrum_survey_manager.get_survey_library(updated_since=now) assert len(surveys) == 1 assert "29333264" == surveys[0].survey_id assert s.is_unchanged(surveys[0]) try: - spectrum_manager.create(s) + spectrum_survey_manager.create(s) except IntegrityError as e: print(e.args) - def test_survey_update(self, settings, spectrum_manager, spectrum_rw): + def test_survey_update( + self, + spectrum_survey_manager: SpectrumSurveyManager, + spectrum_rw: SqlHelper, + spectrum_api_survey_json: dict[str, Any], + ): from generalresearch.models.spectrum.survey import SpectrumSurvey - assert settings.debug, "CRITICAL: Do not run this on production." + assert is_debug(), "CRITICAL: Do not run this on production." now = datetime.now(tz=UTC) spectrum_rw.execute_sql_query( @@ -100,14 +68,13 @@ class TestSpectrumSurvey: """, commit=True, ) - d = copy.deepcopy(example_survey_api_response) - s = SpectrumSurvey.from_api(d) - print(s) + s = SpectrumSurvey.from_api(spectrum_api_survey_json) + assert isinstance(s, SpectrumSurvey) - spectrum_manager.create(s) + spectrum_survey_manager.create(s) s.cpi = Decimal("0.50") - spectrum_manager.update([s]) - surveys = spectrum_manager.get_survey_library(updated_since=now) + spectrum_survey_manager.update([s]) + surveys = spectrum_survey_manager.get_survey_library(updated_since=now) assert len(surveys) == 1 assert "29333264" == surveys[0].survey_id assert Decimal("0.50") == surveys[0].cpi @@ -122,8 +89,8 @@ class TestSpectrumSurvey: s.bid_loi = None s.overall_loi = 1000 s.last_block_loi = 1000 - spectrum_manager.update([s]) - surveys = spectrum_manager.get_survey_library(updated_since=now) + spectrum_survey_manager.update([s]) + surveys = spectrum_survey_manager.get_survey_library(updated_since=now) assert 600 == surveys[0].bid_loi assert 1000 == surveys[0].overall_loi assert 1000 == surveys[0].last_block_loi diff --git a/tests/models/test_currency.py b/tests/models/test_currency.py index 9bc2216..e946126 100644 --- a/tests/models/test_currency.py +++ b/tests/models/test_currency.py @@ -3,6 +3,8 @@ functionality is the same, but pasting here so the tests are in the correct spot... """ +from __future__ import annotations + from decimal import Decimal from random import randint diff --git a/tests/models/test_device.py b/tests/models/test_device.py index bf72c81..8e1251a 100644 --- a/tests/models/test_device.py +++ b/tests/models/test_device.py @@ -1,3 +1,5 @@ +from __future__ import annotations + iphone_ua_string = ( "Mozilla/5.0 (iPhone; CPU iPhone OS 5_1 like Mac OS X) AppleWebKit/534.46 (KHTML, like Gecko) " "Version/5.1 Mobile/9B179 Safari/7534.48.3" @@ -13,10 +15,12 @@ chromebook_ua_string = ( ) +from generalresearch.models import DeviceType +from generalresearch.models.device import parse_device_from_useragent + + class TestDeviceUA: def test_device_ua(self): - from generalresearch.models import DeviceType - from generalresearch.models.device import parse_device_from_useragent assert parse_device_from_useragent(iphone_ua_string) == DeviceType.MOBILE assert parse_device_from_useragent(ipad_ua_string) == DeviceType.TABLET diff --git a/tests/models/test_finance.py b/tests/models/test_finance.py index f84d0b6..72f4f4d 100644 --- a/tests/models/test_finance.py +++ b/tests/models/test_finance.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from collections.abc import Callable from datetime import UTC, datetime, timedelta from itertools import product as iter_product @@ -25,14 +27,13 @@ from generalresearch.models.thl.finance import ( 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 -from test_utils.managers.ledger.conftest import ( - session_with_tx_factory: Callable[..., None], -) fake = Faker() @@ -210,6 +211,8 @@ class TestProductBalanceInitialize: # Confirm the @property computed fields show up in openapi. I don't # know how to do that yet... so this is check to confirm they're # known computed fields for now + + assert isinstance(instance, ProductBalances) computed_fields = list(instance.model_computed_fields.keys()) assert "payout" in computed_fields assert "adjustment" in computed_fields @@ -665,17 +668,18 @@ class TestProductFinanceData: def test_base( self, - product: product: Product, + product: Product, user_factory: Callable[..., User], start: datetime, duration: timedelta, - thl_lm: ThlLedgerManager, + thl_ledger_manager: ThlLedgerManager, + session_with_tx_factory: Callable[..., None], ): # -- Build & Setup # assert ledger_collection.start is None # assert ledger_collection.offset is None - u: User = user_factory(product=product: Product, created=ledger_collection.start) + u: User = user_factory(product=product, created=ledger_collection.start) for item in ledger_collection.items: @@ -699,7 +703,7 @@ class TestProductFinanceData: item_finishes.sort(reverse=True) # -- - account = thl_lm.get_account_or_create_bp_wallet(product=u.product) + account = thl_ledger_manager.get_account_or_create_bp_wallet(product=u.product) ddf = pop_ledger_merge.ddf( force_rr_latest=False, @@ -748,7 +752,7 @@ 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: Callable[..., None], @@ -791,7 +795,7 @@ class TestPOPFinancialData: last_item_finish = item_finishes[0] accounts = [] - for user in users: + for _ in users: account = thl_lm.get_account_or_create_bp_wallet(product=u.product) accounts.append(account) account_ids = [a.uuid for a in accounts] @@ -808,6 +812,7 @@ class TestPOPFinancialData: ("time_idx", "<", last_item_finish), ], ) + df: pd.DataFrame = client_no_amm.compute(collections=ddf, sync=True) df = df.groupby([pd.Grouper(key="time_idx", freq="D"), "account_id"]).sum() @@ -846,16 +851,15 @@ class TestBusinessBalanceData: ledger_collection: LedgerDFCollection, pop_ledger_merge: PopLedgerMerge, user_factory: Callable[..., User], - product: product: Product, + product: Product, create_main_accounts: Callable[..., None], thl_lm: ThlLedgerManager, thl_web_rr: PostgresConfig, delete_df_collection: Callable[..., None], delete_ledger_db: Callable[..., None], session_with_tx_factory: Callable[..., Session], - rm_ledger_collection, + rm_ledger_collection: Callable[..., None], ): - from generalresearch.models.thl.ledger import LedgerAccount delete_ledger_db() create_main_accounts() @@ -863,7 +867,7 @@ class TestBusinessBalanceData: rm_ledger_collection() for _ in range(5): - u: User = user_factory(product=product: Product, created=ledger_collection.start) + u: User = user_factory(product=product, created=ledger_collection.start) for item in ledger_collection.items: item_time = fake.date_time_between( diff --git a/tests/models/thl/question/test_question_info.py b/tests/models/thl/question/test_question_info.py index b619fc3..af8d2b9 100644 --- a/tests/models/thl/question/test_question_info.py +++ b/tests/models/thl/question/test_question_info.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from generalresearch.models.thl.profiling.upk_property import ( ProfilingInfo, UpkProperty, @@ -6,140 +8,9 @@ from generalresearch.models.thl.profiling.upk_property import ( class TestQuestionInfo: - def test_init(self): + def test_init(self, profiling_info_json: str): - s = ( - '[{"property_label": "hispanic", "cardinality": "*", "prop_type": "i", "country_iso": "us", ' - '"property_id": "05170ae296ab49178a075cab2a2073a6", "item_id": "7911ec1468b146ee870951f8ae9cbac1", ' - '"item_label": "panamanian", "gold_standard": 1, "options": [{"id": "c358c11e72c74fa2880358f1d4be85ab", ' - '"label": "not_hispanic"}, {"id": "b1d6c475770849bc8e0200054975dc9c", "label": "yes_hispanic"}, ' - '{"id": "bd1eb44495d84b029e107c188003c2bd", "label": "other_hispanic"}, ' - '{"id": "f290ad5e75bf4f4ea94dc847f57c1bd3", "label": "mexican"}, ' - '{"id": "49f50f2801bd415ea353063bfc02d252", "label": "puerto_rican"}, ' - '{"id": "dcbe005e522f4b10928773926601f8bf", "label": "cuban"}, ' - '{"id": "467ef8ddb7ac4edb88ba9ef817cbb7e9", "label": "salvadoran"}, ' - '{"id": "3c98e7250707403cba2f4dc7b877c963", "label": "dominican"}, ' - '{"id": "981ee77f6d6742609825ef54fea824a8", "label": "guatemalan"}, ' - '{"id": "81c8057b809245a7ae1b8a867ea6c91e", "label": "colombian"}, ' - '{"id": "513656d5f9e249fa955c3b527d483b93", "label": "honduran"}, ' - '{"id": "afc8cddd0c7b4581bea24ccd64db3446", "label": "ecuadorian"}, ' - '{"id": "61f34b36e80747a89d85e1eb17536f84", "label": "argentinian"}, ' - '{"id": "5330cfa681d44aa8ade3a6d0ea198e44", "label": "peruvian"}, ' - '{"id": "e7bceaffd76e486596205d8545019448", "label": "nicaraguan"}, ' - '{"id": "b7bbb2ebf8424714962e6c4f43275985", "label": "spanish"}, ' - '{"id": "8bf539785e7a487892a2f97e52b1932d", "label": "venezuelan"}, ' - '{"id": "7911ec1468b146ee870951f8ae9cbac1", "label": "panamanian"}], "category": [{"id": ' - '"4fd8381d5a1c4409ab007ca254ced084", "label": "Demographic", "path": "/Demographic", ' - '"adwords_vertical_id": null}]}, {"property_label": "ethnic_group", "cardinality": "*", "prop_type": ' - '"i", "country_iso": "us", "property_id": "15070958225d4132b7f6674fcfc979f6", "item_id": ' - '"64b7114cf08143949e3bcc3d00a5d8a0", "item_label": "other_ethnicity", "gold_standard": 1, "options": [{' - '"id": "a72e97f4055e4014a22bee4632cbf573", "label": "caucasians"}, ' - '{"id": "4760353bc0654e46a928ba697b102735", "label": "black_or_african_american"}, ' - '{"id": "20ff0a2969fa4656bbda5c3e0874e63b", "label": "asian"}, ' - '{"id": "107e0a79e6b94b74926c44e70faf3793", "label": "native_hawaiian_or_other_pacific_islander"}, ' - '{"id": "900fa12691d5458c8665bf468f1c98c1", "label": "native_americans"}, ' - '{"id": "64b7114cf08143949e3bcc3d00a5d8a0", "label": "other_ethnicity"}], "category": [{"id": ' - '"4fd8381d5a1c4409ab007ca254ced084", "label": "Demographic", "path": "/Demographic", ' - '"adwords_vertical_id": null}]}, {"property_label": "educational_attainment", "cardinality": "?", ' - '"prop_type": "i", "country_iso": "us", "property_id": "2637783d4b2b4075b93e2a156e16e1d8", "item_id": ' - '"934e7b81d6744a1baa31bbc51f0965d5", "item_label": "other_education", "gold_standard": 1, "options": [{' - '"id": "df35ef9e474b4bf9af520aa86630202d", "label": "3rd_grade_completion"}, ' - '{"id": "83763370a1064bd5ba76d1b68c4b8a23", "label": "8th_grade_completion"}, ' - '{"id": "f0c25a0670c340bc9250099dcce50957", "label": "not_high_school_graduate"}, ' - '{"id": "02ff74c872bd458983a83847e1a9f8fd", "label": "high_school_completion"}, ' - '{"id": "ba8beb807d56441f8fea9b490ed7561c", "label": "vocational_program_completion"}, ' - '{"id": "65373a5f348a410c923e079ddbb58e9b", "label": "some_college_completion"}, ' - '{"id": "2d15d96df85d4cc7b6f58911fdc8d5e2", "label": "associate_academic_degree_completion"}, ' - '{"id": "497b1fedec464151b063cd5367643ffa", "label": "bachelors_degree_completion"}, ' - '{"id": "295133068ac84424ae75e973dc9f2a78", "label": "some_graduate_completion"}, ' - '{"id": "e64f874faeff4062a5aa72ac483b4b9f", "label": "masters_degree_completion"}, ' - '{"id": "cbaec19a636d476385fb8e7842b044f5", "label": "doctorate_degree_completion"}, ' - '{"id": "934e7b81d6744a1baa31bbc51f0965d5", "label": "other_education"}], "category": [{"id": ' - '"4fd8381d5a1c4409ab007ca254ced084", "label": "Demographic", "path": "/Demographic", ' - '"adwords_vertical_id": null}]}, {"property_label": "household_spoken_language", "cardinality": "*", ' - '"prop_type": "i", "country_iso": "us", "property_id": "5a844571073d482a96853a0594859a51", "item_id": ' - '"62b39c1de141422896ad4ab3c4318209", "item_label": "dut", "gold_standard": 1, "options": [{"id": ' - '"f65cd57b79d14f0f8460761ce41ec173", "label": "ara"}, {"id": "6d49de1f8f394216821310abd29392d9", ' - '"label": "zho"}, {"id": "be6dc23c2bf34c3f81e96ddace22800d", "label": "eng"}, ' - '{"id": "ddc81f28752d47a3b1c1f3b8b01a9b07", "label": "fre"}, {"id": "2dbb67b29bd34e0eb630b1b8385542ca", ' - '"label": "ger"}, {"id": "a747f96952fc4b9d97edeeee5120091b", "label": "hat"}, ' - '{"id": "7144b04a3219433baac86273677551fa", "label": "hin"}, {"id": "e07ff3e82c7149eaab7ea2b39ee6a6dc", ' - '"label": "ita"}, {"id": "b681eff81975432ebfb9f5cc22dedaa3", "label": "jpn"}, ' - '{"id": "5cb20440a8f64c9ca62fb49c1e80cdef", "label": "kor"}, {"id": "171c4b77d4204bc6ac0c2b81e38a10ff", ' - '"label": "pan"}, {"id": "8c3ec18e6b6c4a55a00dd6052e8e84fb", "label": "pol"}, ' - '{"id": "3ce074d81d384dd5b96f1fb48f87bf01", "label": "por"}, {"id": "6138dc951990458fa88a666f6ddd907b", ' - '"label": "rus"}, {"id": "e66e5ecc07df4ebaa546e0b436f034bd", "label": "spa"}, ' - '{"id": "5a981b3d2f0d402a96dd2d0392ec2fcb", "label": "tgl"}, {"id": "b446251bd211403487806c4d0a904981", ' - '"label": "vie"}, {"id": "92fb3ee337374e2db875fb23f52eed46", "label": "xxx"}, ' - '{"id": "8b1f590f12f24cc1924d7bdcbe82081e", "label": "ind"}, {"id": "bf3f4be556a34ff4b836420149fd2037", ' - '"label": "tur"}, {"id": "87ca815c43ba4e7f98cbca98821aa508", "label": "zul"}, ' - '{"id": "0adbf915a7a64d67a87bb3ce5d39ca54", "label": "may"}, {"id": "62b39c1de141422896ad4ab3c4318209", ' - '"label": "dut"}], "category": [{"id": "4fd8381d5a1c4409ab007ca254ced084", "label": "Demographic", ' - '"path": "/Demographic", "adwords_vertical_id": null}]}, {"property_label": "gender", "cardinality": ' - '"?", "prop_type": "i", "country_iso": "us", "property_id": "73175402104741549f21de2071556cd7", ' - '"item_id": "093593e316344cd3a0ac73669fca8048", "item_label": "other_gender", "gold_standard": 1, ' - '"options": [{"id": "b9fc5ea07f3a4252a792fd4a49e7b52b", "label": "male"}, ' - '{"id": "9fdb8e5e18474a0b84a0262c21e17b56", "label": "female"}, ' - '{"id": "093593e316344cd3a0ac73669fca8048", "label": "other_gender"}], "category": [{"id": ' - '"4fd8381d5a1c4409ab007ca254ced084", "label": "Demographic", "path": "/Demographic", ' - '"adwords_vertical_id": null}]}, {"property_label": "age_in_years", "cardinality": "?", "prop_type": ' - '"n", "country_iso": "us", "property_id": "94f7379437874076b345d76642d4ce6d", "item_id": null, ' - '"item_label": null, "gold_standard": 1, "category": [{"id": "4fd8381d5a1c4409ab007ca254ced084", ' - '"label": "Demographic", "path": "/Demographic", "adwords_vertical_id": null}]}, {"property_label": ' - '"children_age_gender", "cardinality": "*", "prop_type": "i", "country_iso": "us", "property_id": ' - '"e926142fcea94b9cbbe13dc7891e1e7f", "item_id": "b7b8074e95334b008e8958ccb0a204f1", "item_label": ' - '"female_18", "gold_standard": 1, "options": [{"id": "16a6448ec24c48d4993d78ebee33f9b4", ' - '"label": "male_under_1"}, {"id": "809c04cb2e3b4a3bbd8077ab62cdc220", "label": "female_under_1"}, ' - '{"id": "295e05bb6a0843bc998890b24c99841e", "label": "no_children"}, ' - '{"id": "142cb948d98c4ae8b0ef2ef10978e023", "label": "male_0"}, ' - '{"id": "5a5c1b0e9abc48a98b3bc5f817d6e9d0", "label": "male_1"}, ' - '{"id": "286b1a9afb884bdfb676dbb855479d1e", "label": "male_2"}, ' - '{"id": "942ca3cda699453093df8cbabb890607", "label": "male_3"}, ' - '{"id": "995818d432f643ec8dd17e0809b24b56", "label": "male_4"}, ' - '{"id": "f38f8b57f25f4cdea0f270297a1e7a5c", "label": "male_5"}, ' - '{"id": "975df709e6d140d1a470db35023c432d", "label": "male_6"}, ' - '{"id": "f60bd89bbe0f4e92b90bccbc500467c2", "label": "male_7"}, ' - '{"id": "6714ceb3ed5042c0b605f00b06814207", "label": "male_8"}, ' - '{"id": "c03c2f8271d443cf9df380e84b4dea4c", "label": "male_9"}, ' - '{"id": "11690ee0f5a54cb794f7ddd010d74fa2", "label": "male_10"}, ' - '{"id": "17bef9a9d14b4197b2c5609fa94b0642", "label": "male_11"}, ' - '{"id": "e79c8338fe28454f89ccc78daf6f409a", "label": "male_12"}, ' - '{"id": "3a4f87acb3fa41f4ae08dfe2858238c1", "label": "male_13"}, ' - '{"id": "36ffb79d8b7840a7a8cb8d63bbc8df59", "label": "male_14"}, ' - '{"id": "1401a508f9664347aee927f6ec5b0a40", "label": "male_15"}, ' - '{"id": "6e0943c5ec4a4f75869eb195e3eafa50", "label": "male_16"}, ' - '{"id": "47d4b27b7b5242758a9fff13d3d324cf", "label": "male_17"}, ' - '{"id": "9ce886459dd44c9395eb77e1386ab181", "label": "female_0"}, ' - '{"id": "6499ccbf990d4be5b686aec1c7353fd8", "label": "female_1"}, ' - '{"id": "d85ceaa39f6d492abfc8da49acfd14f2", "label": "female_2"}, ' - '{"id": "18edb45c138e451d8cb428aefbb80f9c", "label": "female_3"}, ' - '{"id": "bac6f006ed9f4ccf85f48e91e99fdfd1", "label": "female_4"}, ' - '{"id": "5a6a1a8ad00c4ce8be52dcb267b034ff", "label": "female_5"}, ' - '{"id": "6bff0acbf6364c94ad89507bcd5f4f45", "label": "female_6"}, ' - '{"id": "d0d56a0a6b6f4516a366a2ce139b4411", "label": "female_7"}, ' - '{"id": "bda6028468044b659843e2bef4db2175", "label": "female_8"}, ' - '{"id": "dbb6d50325464032b456357b1a6e5e9c", "label": "female_9"}, ' - '{"id": "b87a93d7dc1348edac5e771684d63fb8", "label": "female_10"}, ' - '{"id": "11449d0d98f14e27ba47de40b18921d7", "label": "female_11"}, ' - '{"id": "16156501e97b4263962cbbb743840292", "label": "female_12"}, ' - '{"id": "04ee971c89a345cc8141a45bce96050c", "label": "female_13"}, ' - '{"id": "e818d310bfbc4faba4355e5d2ed49d4f", "label": "female_14"}, ' - '{"id": "440d25e078924ba0973163153c417ed6", "label": "female_15"}, ' - '{"id": "78ff804cc9b441c5a524bd91e3d1f8bf", "label": "female_16"}, ' - '{"id": "4b04d804d7d84786b2b1c22e4ed440f5", "label": "female_17"}, ' - '{"id": "28bc848cd3ff44c3893c76bfc9bc0c4e", "label": "male_18"}, ' - '{"id": "b7b8074e95334b008e8958ccb0a204f1", "label": "female_18"}], "category": [{"id": ' - '"e18ba6e9d51e482cbb19acf2e6f505ce", "label": "Parenting", "path": "/People & Society/Family & ' - 'Relationships/Family/Parenting", "adwords_vertical_id": "58"}]}, {"property_label": "home_postal_code", ' - '"cardinality": "?", "prop_type": "x", "country_iso": "us", "property_id": ' - '"f3b32ebe78014fbeb1ed6ff77d6338bf", "item_id": null, "item_label": null, "gold_standard": 1, ' - '"category": [{"id": "4fd8381d5a1c4409ab007ca254ced084", "label": "Demographic", "path": "/Demographic", ' - '"adwords_vertical_id": null}]}, {"property_label": "household_income", "cardinality": "?", "prop_type": ' - '"n", "country_iso": "us", "property_id": "ff5b1d4501d5478f98de8c90ef996ac1", "item_id": null, ' - '"item_label": null, "gold_standard": 1, "category": [{"id": "4fd8381d5a1c4409ab007ca254ced084", ' - '"label": "Demographic", "path": "/Demographic", "adwords_vertical_id": null}]}]' - ) - instance_list = ProfilingInfo.validate_json(s) + instance_list = ProfilingInfo.validate_json(profiling_info_json) assert isinstance(instance_list, list) for i in instance_list: diff --git a/tests/models/thl/question/test_user_info.py b/tests/models/thl/question/test_user_info.py index 0bbbc78..5410d35 100644 --- a/tests/models/thl/question/test_user_info.py +++ b/tests/models/thl/question/test_user_info.py @@ -1,32 +1,11 @@ +from __future__ import annotations + from generalresearch.models.thl.profiling.user_info import UserInfo class TestUserInfo: - def test_init(self): + def test_init(self, profiling_user_info_json: str): - s = ( - '{"user_profile_knowledge": [], "marketplace_profile_knowledge": [{"source": "d", "question_id": ' - '"1", "answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "pr", ' - '"question_id": "3", "answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": ' - '"h", "question_id": "60", "answer": ["58"], "created": "2023-11-07T16:41:05.234096Z"}, ' - '{"source": "c", "question_id": "43", "answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, ' - '{"source": "s", "question_id": "211", "answer": ["111"], "created": ' - '"2023-11-07T16:41:05.234096Z"}, {"source": "s", "question_id": "1843", "answer": ["111"], ' - '"created": "2023-11-07T16:41:05.234096Z"}, {"source": "h", "question_id": "13959", "answer": [' - '"244155"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "c", "question_id": "33092", ' - '"answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "c", "question_id": "gender", ' - '"answer": ["10682"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "e", "question_id": ' - '"gender", "answer": ["male"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "f", ' - '"question_id": "gender", "answer": ["male"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": ' - '"i", "question_id": "gender", "answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, ' - '{"source": "c", "question_id": "137510", "answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, ' - '{"source": "m", "question_id": "gender", "answer": ["1"], "created": ' - '"2023-11-07T16:41:05.234096Z"}, {"source": "o", "question_id": "gender", "answer": ["male"], ' - '"created": "2023-11-07T16:41:05.234096Z"}, {"source": "c", "question_id": "gender_plus", "answer": [' - '"7657644"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "i", "question_id": ' - '"gender_plus", "answer": ["1"], "created": "2023-11-07T16:41:05.234096Z"}, {"source": "c", ' - '"question_id": "income_level", "answer": ["9071"], "created": "2023-11-07T16:41:05.234096Z"}]}' - ) - instance = UserInfo.model_validate_json(s) + instance = UserInfo.model_validate_json(profiling_user_info_json) assert isinstance(instance, UserInfo) diff --git a/tests/models/thl/test_adjustments.py b/tests/models/thl/test_adjustments.py index 91e5316..c5b3f6b 100644 --- a/tests/models/thl/test_adjustments.py +++ b/tests/models/thl/test_adjustments.py @@ -1,9 +1,13 @@ +from __future__ import annotations + from collections.abc import Callable from datetime import UTC, datetime, timedelta from decimal import Decimal import pytest +from generalresearch.managers.thl.session import SessionManager +from generalresearch.managers.thl.wall import WallManager from generalresearch.models import Source from generalresearch.models.thl.product import Product from generalresearch.models.thl.session import ( @@ -30,7 +34,7 @@ class TestProductAdjustments: @pytest.mark.parametrize("payout", [".6", "1", "1.8", "2", "500.0000"]) def test_determine_bp_payment_no_rounding( - self, product_factory: Callable[..., Product], payout + self, product_factory: Callable[..., Product], payout: str ): p1 = product_factory(commission_pct=Decimal("0.05")) res = p1.determine_bp_payment(thl_net=Decimal(payout)) @@ -39,7 +43,7 @@ class TestProductAdjustments: @pytest.mark.parametrize("payout", [".01", ".05", ".5"]) def test_determine_bp_payment_rounding( - self, product_factory: Callable[..., Product], payout + self, product_factory: Callable[..., Product], payout: str ): p1 = product_factory(commission_pct=Decimal("0.05")) res = p1.determine_bp_payment(thl_net=Decimal(payout)) @@ -73,7 +77,10 @@ class TestSessionAdjustments: class TestAdjustments: def test_finish_with_status( - self, session_factory: Callable[..., Session], user: User, session_manager + self, + session_factory: Callable[..., Session], + user: User, + session_manager: SessionManager, ): # Completed Session with 2 wall events s1 = session_factory( @@ -85,6 +92,7 @@ class TestAdjustments: ) status, status_code_1 = s1.determine_session_status() + assert isinstance(user.product, Product) payout = user.product.determine_bp_payment(Decimal(1)) session_manager.finish_with_status( session=s1, @@ -97,7 +105,10 @@ class TestAdjustments: assert Decimal("0.95") == payout def test_never_adjusted( - self, session_factory: Callable[..., Session], user: User, session_manager + self, + session_factory: Callable[..., Session], + user: User, + session_manager: SessionManager, ): s1 = session_factory( user=user, @@ -130,8 +141,8 @@ class TestAdjustments: self, session_factory: Callable[..., Session], user: User, - session_manager, - wall_manager, + session_manager: SessionManager, + wall_manager: WallManager, ): # Completed Session with 2 wall events s1 = session_factory( @@ -174,13 +185,14 @@ class TestAdjustments: # Because the Product doesn't have the Wallet mode enabled, the # user_payout fields should always be None + assert isinstance(user.product, Product) assert not user.product.user_wallet_config.enabled assert s1.adjusted_user_payout is None def test_adjustment_session_values( self, - wall_manager, - session_manager, + wall_manager: WallManager, + session_manager: SessionManager, session_factory: Callable[..., Session], user: User, ): @@ -218,13 +230,14 @@ class TestAdjustments: # Because the Product doesn't have the Wallet mode enabled, the # user_payout fields should always be None + assert isinstance(user.product, Product) assert not user.product.user_wallet_config.enabled assert s1.adjusted_user_payout is None def test_double_adjustment_session_values( self, - wall_manager, - session_manager, + wall_manager: WallManager, + session_manager: SessionManager, session_factory: Callable[..., Session], user: User, ): @@ -276,8 +289,8 @@ class TestAdjustments: def test_double_adjustment_sm_vs_db_values( self, - wall_manager, - session_manager, + wall_manager: WallManager, + session_manager: SessionManager, session_factory: Callable[..., Session], user: User, ): @@ -343,8 +356,8 @@ class TestAdjustments: def test_double_adjustment_double_completes( self, - wall_manager, - session_manager, + wall_manager: WallManager, + session_manager: SessionManager, session_factory: Callable[..., Session], user: User, ): @@ -419,8 +432,8 @@ class TestAdjustments: self, session_factory: Callable[..., Session], user: User, - session_manager, - wall_manager, + session_manager: SessionManager, + wall_manager: WallManager, utc_hour_ago: datetime, ): s1 = session_factory( @@ -435,6 +448,7 @@ class TestAdjustments: assert status == Status.COMPLETE thl_net = Decimal(sum(w.cpi for w in s1.wall_events if w.is_visible_complete())) + assert isinstance(user.product, Product) payout = user.product.determine_bp_payment(thl_net=thl_net) session_manager.finish_with_status( @@ -459,7 +473,7 @@ class TestAdjustments: assert Status.FAIL == new_status assert Decimal(0) == new_payout - assert isinstance(user.product: Product, Product) + assert isinstance(user.product, Product) assert not user.product.user_wallet_config.enabled assert new_user_payout is None @@ -560,6 +574,7 @@ class TestAdjustments: new_status, new_payout, new_user_payout = s1.determine_new_status_and_payouts() assert Status.COMPLETE == new_status assert Decimal("0.95") == new_payout + assert isinstance(user.product, Product) assert not user.product.user_wallet_config.enabled # assert Decimal("0.48") == new_user_payout assert new_user_payout is None @@ -588,6 +603,7 @@ class TestAdjustments: status, status_code_1 = s1.determine_session_status() thl_net = Decimal(sum(w.cpi for w in s1.wall_events if w.is_visible_complete())) + assert isinstance(user.product, Product) payout = user.product.determine_bp_payment(thl_net=thl_net) s1.update( status=status, @@ -624,7 +640,10 @@ class TestAdjustments: assert s1.adjusted_user_payout is None def test_complete_to_fail_to_complete_adj1( - self, user, session_factory, utc_hour_ago + self, + user: User, + session_factory: Callable[..., Session], + utc_hour_ago: datetime, ): # Same as test_complete_to_fail_to_complete_adj but in opposite order s1 = session_factory( @@ -640,6 +659,7 @@ class TestAdjustments: status, status_code_1 = s1.determine_session_status() thl_net = Decimal(sum(w.cpi for w in s1.wall_events if w.is_visible_complete())) + assert isinstance(user.product, Product) payout = user.product.determine_bp_payment(thl_net) s1.update( status=status, @@ -658,6 +678,7 @@ class TestAdjustments: s1.adjust_status() assert SessionAdjustedStatus.ADJUSTED_TO_FAIL == s1.adjusted_status assert Decimal(0) == s1.adjusted_payout + assert isinstance(user.product, Product) assert not user.product.user_wallet_config.enabled # assert Decimal(0) == s.adjusted_user_payout assert s1.adjusted_user_payout is None @@ -702,6 +723,7 @@ class TestAdjustments: s1.adjust_status() assert SessionAdjustedStatus.ADJUSTED_TO_COMPLETE == s1.adjusted_status assert Decimal("1.90") == s1.adjusted_payout + assert isinstance(user.product, Product) assert not user.product.user_wallet_config.enabled # assert Decimal("0.95") == s1.adjusted_user_payout assert s1.adjusted_user_payout is None diff --git a/tests/models/thl/test_bucket.py b/tests/models/thl/test_bucket.py index 0aa5843..8d2f728 100644 --- a/tests/models/thl/test_bucket.py +++ b/tests/models/thl/test_bucket.py @@ -1,14 +1,17 @@ +from __future__ import annotations + from datetime import timedelta from decimal import Decimal import pytest from pydantic import ValidationError +from generalresearch.models.legacy.bucket import Bucket + class TestBucket: def test_raises_payout(self): - from generalresearch.models.legacy.bucket import Bucket with pytest.raises(expected_exception=ValidationError) as e: Bucket(user_payout_min=123) @@ -27,7 +30,6 @@ class TestBucket: assert "user_payout_min should be <= user_payout_max" in str(e.value) def test_raises_loi(self): - from generalresearch.models.legacy.bucket import Bucket with pytest.raises(expected_exception=ValidationError) as e: Bucket(loi_min=123) @@ -63,7 +65,6 @@ class TestBucket: assert "loi_q1 should be <= loi_q2" in str(e.value) def test_parse_1(self): - from generalresearch.models.legacy.bucket import Bucket b1 = Bucket.parse_from_offerwall({"payout": {"min": 123}}) b_exp = Bucket( @@ -180,7 +181,6 @@ class TestBucket: assert b_exp == b4 def test_parse_3(self): - from generalresearch.models.legacy.bucket import Bucket b1 = Bucket.parse_from_offerwall({"payout": 123}) b_exp = Bucket( diff --git a/tests/models/thl/test_buyer.py b/tests/models/thl/test_buyer.py index eebb828..02093e2 100644 --- a/tests/models/thl/test_buyer.py +++ b/tests/models/thl/test_buyer.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from generalresearch.models 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 acb501c..e1053f4 100644 --- a/tests/models/thl/test_contest/test_contest.py +++ b/tests/models/thl/test_contest/test_contest.py @@ -1,3 +1,5 @@ +from __future__ import annotations + 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 5bab060..99cfb37 100644 --- a/tests/models/thl/test_contest/test_leaderboard_contest.py +++ b/tests/models/thl/test_contest/test_leaderboard_contest.py @@ -1,10 +1,14 @@ +from __future__ import annotations + from datetime import UTC from uuid import uuid4 import pytest +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, @@ -18,6 +22,7 @@ from generalresearch.models.thl.contest.utils import ( ) 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 @@ -25,7 +30,7 @@ class TestLeaderboardContest(TestContest): @pytest.fixture def leaderboard_contest( - self, product: product: Product, thl_redis, user_manager + self, product: Product, thl_redis: Redis, user_manager: UserManager ) -> LeaderboardContest: board_key = f"leaderboard:{product.uuid}:us:weekly:2025-05-26:complete_count" @@ -63,7 +68,13 @@ class TestLeaderboardContest(TestContest): c._user_manager = user_manager return c - def test_init(self, leaderboard_contest, thl_redis, user_1, user_2): + def test_init( + self, + leaderboard_contest: LeaderboardContest, + thl_redis: Redis, + user_1: User, + user_2: User, + ): model = leaderboard_contest.leaderboard_model assert leaderboard_contest.end_condition.ends_at is not None @@ -83,7 +94,14 @@ class TestLeaderboardContest(TestContest): lb = leaderboard_contest.get_leaderboard() print(lb) - def test_win(self, leaderboard_contest, thl_redis, user_1, user_2, user_3): + def test_win( + self, + leaderboard_contest: LeaderboardContest, + thl_redis: Redis, + user_1: User, + user_2: User, + user_3: User, + ): model = leaderboard_contest.leaderboard_model lbm = LeaderboardManager( redis_client=thl_redis, @@ -102,10 +120,13 @@ class TestLeaderboardContest(TestContest): lbm.hit_complete_count(product_user_id=user_3.product_user_id) leaderboard_contest.end_contest() + assert isinstance(leaderboard_contest.all_winners, list) assert len(leaderboard_contest.all_winners) == 3 # Prizes are $15, $10, $5. user 2 and 3 ties for 2nd place, so they split (10 + 5) assert leaderboard_contest.all_winners[0].awarded_cash_amount == USDCent(15_00) + + assert isinstance(leaderboard_contest.all_winners[0].user, User) assert ( leaderboard_contest.all_winners[0].user.product_user_id == user_1.product_user_id diff --git a/tests/models/thl/test_contest/test_raffle_contest.py b/tests/models/thl/test_contest/test_raffle_contest.py index f85ba75..8812cb3 100644 --- a/tests/models/thl/test_contest/test_raffle_contest.py +++ b/tests/models/thl/test_contest/test_raffle_contest.py @@ -1,4 +1,7 @@ +from __future__ import annotations + from collections import Counter +from datetime import datetime from uuid import uuid4 import pytest @@ -19,6 +22,7 @@ from generalresearch.models.thl.contest.definitions import ( ) 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 @@ -42,7 +46,9 @@ class TestRaffleContest(TestContest): ) @pytest.fixture(scope="function") - def ended_raffle_contest(self, raffle_contest, utc_now) -> RaffleContest: + def ended_raffle_contest( + self, raffle_contest: RaffleContest, utc_now: datetime + ) -> RaffleContest: # Fake ending the contest raffle_contest = raffle_contest.model_copy() raffle_contest.update( @@ -55,7 +61,7 @@ class TestRaffleContest(TestContest): class TestRaffleContestUserView(TestRaffleContest): - def test_user_view(self, raffle_contest, user): + def test_user_view(self, raffle_contest: RaffleContest, user: User): from generalresearch.models.thl.contest.raffle import RaffleUserView data = { @@ -78,7 +84,7 @@ class TestRaffleContestUserView(TestRaffleContest): assert res["current_win_probability"] == approx(0.0099, rel=0.001) assert res["projected_win_probability"] == approx(0.0099, rel=0.001) - def test_win_pct(self, raffle_contest, user): + def test_win_pct(self, raffle_contest: RaffleContest, user: User): from generalresearch.models.thl.contest.raffle import RaffleUserView data = { @@ -124,7 +130,9 @@ class TestRaffleContestUserView(TestRaffleContest): class TestRaffleContestWinners(TestRaffleContest): - def test_winners_1_prize(self, ended_raffle_contest, user_1, user_2, user_3): + def test_winners_1_prize( + self, ended_raffle_contest, user_1: User, user_2: User, user_3: User + ): ended_raffle_contest.entries = [ ContestEntry( user=user_1, @@ -160,7 +168,13 @@ class TestRaffleContestWinners(TestRaffleContest): assert c[user_2.user_id] == approx(10000 * 2 / 6, rel=0.1) assert c[user_3.user_id] == approx(10000 * 3 / 6, rel=0.1) - def test_winners_2_prizes(self, ended_raffle_contest, user_1, user_2, user_3): + def test_winners_2_prizes( + self, + ended_raffle_contest: RaffleContest, + user_1: User, + user_2: User, + user_3: User, + ): ended_raffle_contest.prizes.append( ContestPrize( name="iPod 64GB Black", @@ -193,7 +207,9 @@ class TestRaffleContestWinners(TestRaffleContest): # Same user assert all(w.user.user_id == user_1.user_id for w in winners) - def test_winners_2_prizes_1_entry(self, ended_raffle_contest, user_3): + def test_winners_2_prizes_1_entry( + self, ended_raffle_contest: RaffleContest, user_3: User + ): ended_raffle_contest.prizes = [ ContestPrize( name="iPod 64GB White", @@ -218,7 +234,9 @@ class TestRaffleContestWinners(TestRaffleContest): winners = ended_raffle_contest.select_winners() assert len(winners) == 1 - def test_winners_2_prizes_1_entry_2_pennies(self, ended_raffle_contest, user_3): + def test_winners_2_prizes_1_entry_2_pennies( + self, ended_raffle_contest: RaffleContest, user_3: User + ): ended_raffle_contest.prizes = [ ContestPrize( name="iPod 64GB White", @@ -243,7 +261,12 @@ class TestRaffleContestWinners(TestRaffleContest): assert len(winners) == 2 def test_winners_3_prizes_3_entries( - self, ended_raffle_contest, product: Product, user_1, user_2, user_3 + self, + ended_raffle_contest: RaffleContest, + product: Product, + user_1: User, + user_2: User, + user_3: User, ): ended_raffle_contest.prizes = [ ContestPrize( diff --git a/tests/models/thl/test_ledger.py b/tests/models/thl/test_ledger.py index 7066180..7c48dbd 100644 --- a/tests/models/thl/test_ledger.py +++ b/tests/models/thl/test_ledger.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from datetime import UTC, datetime from uuid import uuid4 diff --git a/tests/models/thl/test_marketplace_condition.py b/tests/models/thl/test_marketplace_condition.py index 8a4b25c..1dd25e8 100644 --- a/tests/models/thl/test_marketplace_condition.py +++ b/tests/models/thl/test_marketplace_condition.py @@ -1,15 +1,18 @@ +from __future__ import annotations + import pytest from pydantic import ValidationError +from generalresearch.models import LogicalOperator +from generalresearch.models.thl.survey.condition import ( + ConditionValueType, + MarketplaceCondition, +) + class TestMarketplaceCondition: def test_list_or(self): - from generalresearch.models import LogicalOperator - from generalresearch.models.thl.survey.condition import ( - ConditionValueType, - MarketplaceCondition, - ) user_qas = {"q1": {"a2"}} c = MarketplaceCondition( @@ -46,11 +49,6 @@ class TestMarketplaceCondition: assert c.evaluate_criterion(user_qas) is None def test_list_or_negate(self): - from generalresearch.models import LogicalOperator - from generalresearch.models.thl.survey.condition import ( - ConditionValueType, - MarketplaceCondition, - ) user_qas = {"q1": {"a2"}} c = MarketplaceCondition( @@ -87,11 +85,6 @@ class TestMarketplaceCondition: assert c.evaluate_criterion(user_qas) is None def test_list_and(self): - from generalresearch.models import LogicalOperator - from generalresearch.models.thl.survey.condition import ( - ConditionValueType, - MarketplaceCondition, - ) user_qas = {"q1": {"a1", "a2"}} c = MarketplaceCondition( @@ -178,11 +171,6 @@ class TestMarketplaceCondition: assert c.evaluate_criterion(user_qas) is None def test_ranges(self): - from generalresearch.models import LogicalOperator - from generalresearch.models.thl.survey.condition import ( - ConditionValueType, - MarketplaceCondition, - ) user_qas = {"q1": {"2", "50"}} c = MarketplaceCondition( @@ -245,12 +233,6 @@ class TestMarketplaceCondition: ) def test_ranges_to_list(self): - from generalresearch.models import LogicalOperator - from generalresearch.models.thl.survey.condition import ( - ConditionValueType, - MarketplaceCondition, - ) - user_qas = {"q1": {"2", "50"}} MarketplaceCondition._CONVERT_LIST_TO_RANGE = ["q1"] c = MarketplaceCondition( @@ -309,10 +291,6 @@ class TestMarketplaceCondition: assert not c.evaluate_criterion({"q1": {"50"}}) def test_answered(self): - from generalresearch.models.thl.survey.condition import ( - ConditionValueType, - MarketplaceCondition, - ) user_qas = {"q1": {"a2"}} c = MarketplaceCondition( diff --git a/tests/models/thl/test_payout.py b/tests/models/thl/test_payout.py index dd0065c..daf1bd7 100644 --- a/tests/models/thl/test_payout.py +++ b/tests/models/thl/test_payout.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from uuid import uuid4 import pytest @@ -5,7 +7,7 @@ from pydantic import ValidationError from generalresearch.currency import USDCent from generalresearch.models.gr import Team -from generalresearch.models.gr.business import business: Business, BusinessAddress, BusinessType +from generalresearch.models.gr.business import Business, BusinessAddress, BusinessType from generalresearch.models.thl.payout import ( BrokerageProductPayoutEvent, BusinessPayoutEvent, diff --git a/tests/models/thl/test_payout_format.py b/tests/models/thl/test_payout_format.py index 83fde25..fe7aea5 100644 --- a/tests/models/thl/test_payout_format.py +++ b/tests/models/thl/test_payout_format.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import pytest from pydantic import BaseModel diff --git a/tests/models/thl/test_product.py b/tests/models/thl/test_product.py index b7ee654..adf276d 100644 --- a/tests/models/thl/test_product.py +++ b/tests/models/thl/test_product.py @@ -13,22 +13,26 @@ 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.thl.finance import ProductBalances -from generalresearch.models.thl.product import ( +from generalresearch.models.thl.payout import ( BrokerageProductPayoutEvent, +) +from generalresearch.models.thl.product import ( BrokerageProductPayoutEventManager, IntegrationMode, PayoutConfig, PayoutTransformation, PayoutTransformationPercentArgs, - product: Product, + Product, ProfilingConfig, SourceConfig, SourcesConfig, @@ -37,6 +41,7 @@ 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 class TestProduct: @@ -156,6 +161,11 @@ class TestProduct: assert ( "payout_transformation_percent" == p.payout_config.payout_transformation.f ) + + assert isinstance( + p.payout_config.payout_transformation.kwargs, + PayoutTransformationPercentArgs, + ) assert 0.5 == p.payout_config.payout_transformation.kwargs.pct assert ( Decimal("0.10") == p.payout_config.payout_transformation.kwargs.min_payout @@ -287,10 +297,10 @@ class TestProduct: p.profiling_config = ProfilingConfig(max_questions=1) assert p.profiling_config.max_questions == 1 - def test_bp_account(self, product: Product, thl_lm): + def test_bp_account(self, product: Product, thl_ledger_manager: ThlLedgerManager): assert product.bp_account is None - product.prefetch_bp_account(thl_lm=thl_lm) + product.prefetch_bp_account(thl_lm=thl_ledger_manager) from generalresearch.models.thl.ledger import LedgerAccount @@ -391,7 +401,7 @@ class TestGlobalProduct: random_product = uuid4().hex random_team = uuid4().hex res = instance.sources_config.get_policies_for( - product_id=random_product: Product, team_id=random_team + product_id=random_product, team_id=random_team ) assert res == s.global_scoped_policies_dict @@ -598,7 +608,7 @@ class TestProductFinancials: def test_balance( self, - business: business: Business, + business: Business, product_factory: Callable[..., Product], user_factory: Callable[..., User], mnt_filepath: GRLDatasets, @@ -610,7 +620,7 @@ class TestProductFinancials: delete_ledger_db: Callable[..., None], create_main_accounts: Callable[..., None], client_no_amm: DaskClient, - ledger_collection, + ledger_collection: LedgerDFCollection, pop_ledger_merge: PopLedgerMerge, delete_df_collection: Callable[..., None], ): @@ -781,20 +791,20 @@ class TestProductBalance: def test_inconsistent( self, - product: product: Product, + product: Product, mnt_filepath: GRLDatasets, thl_lm: ThlLedgerManager, client_no_amm: DaskClient, delete_ledger_db: Callable[..., None], create_main_accounts: Callable[..., None], delete_df_collection: Callable[..., None], - ledger_collection, + ledger_collection: LedgerDFCollection, user_factory: Callable[..., User], session_with_tx_factory: Callable[..., Session], - pop_ledger_merge, + pop_ledger_merge: PopLedgerMerge, start: datetime, - bp_payout_factory, - payout_event_manager, + bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], + payout_event_manager: PayoutEventManager, ): # Now let's load it up and actually test some things delete_ledger_db() @@ -815,7 +825,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,20 +843,20 @@ 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: Callable[..., None], create_main_accounts: Callable[..., None], delete_df_collection: Callable[..., None], - ledger_collection, + ledger_collection: LedgerDFCollection, user_factory: Callable[..., User], session_with_tx_factory: Callable[..., None], pop_ledger_merge: PopLedgerMerge, start: datetime, - bp_payout_factory, - payout_event_manager, + bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], + payout_event_manager: PayoutEventManager, ): # This is very similar to the test_complete_payout_pq_inconsistent # test, however this time we're only going to assign the payout @@ -874,7 +884,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,14 +914,14 @@ class TestProductPOPFinancial: def test_base( self, - product: product: Product, + product: Product, mnt_filepath: GRLDatasets, thl_lm: ThlLedgerManager, client_no_amm: DaskClient, delete_ledger_db: Callable[..., None], create_main_accounts: Callable[..., None], delete_df_collection: Callable[..., None], - ledger_collection, + ledger_collection: LedgerDFCollection, user_factory: Callable[..., User], session_with_tx_factory: Callable[..., None], pop_ledger_merge: PopLedgerMerge, @@ -977,16 +987,16 @@ class TestProductCache: def test_basic( self, - product: product: Product, - mnt_filepath, - thl_lm, + product: Product, + mnt_filepath: GRLDatasets, + thl_ledger_manager: ThlLedgerManager, client_no_amm: DaskClient, thl_redis_config: RedisConfig, - brokerage_product_payout_event_manager, + brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager, delete_ledger_db: Callable[..., None], create_main_accounts: Callable[..., None], delete_df_collection: Callable[..., None], - ledger_collection, + ledger_collection: LedgerDFCollection, user_factory: Callable[..., User], session_with_tx_factory: Callable[..., None], pop_ledger_merge: PopLedgerMerge, @@ -1007,7 +1017,7 @@ class TestProductCache: ds=mnt_filepath, client=client_no_amm, bp_pem=brokerage_product_payout_event_manager, - redis_config=thl_redis_config: RedisConfig, + redis_config=thl_redis_config, ) from generalresearch.models.thl.product import Product @@ -1029,7 +1039,7 @@ class TestProductCache: ds=mnt_filepath, client=client_no_amm, bp_pem=brokerage_product_payout_event_manager, - redis_config=thl_redis_config: RedisConfig, + redis_config=thl_redis_config, ) # Fetch from cache and assert the instance loaded from redis @@ -1048,23 +1058,23 @@ class TestProductCache: def test_neg_balance_cache( self, - product: product: Product, + product: Product, mnt_filepath: GRLDatasets, - thl_lm, + thl_ledger_manager: ThlLedgerManager, client_no_amm: DaskClient, thl_redis_config: RedisConfig, - brokerage_product_payout_event_manager, + brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager, delete_ledger_db: Callable[..., None], create_main_accounts: Callable[..., None], delete_df_collection: Callable[..., None], - ledger_collection, + ledger_collection: LedgerDFCollection, user_factory: Callable[..., User], session_with_tx_factory: Callable[..., None], pop_ledger_merge: PopLedgerMerge, start: datetime, - bp_payout_factory, - payout_event_manager, - adj_to_fail_with_tx_factory, + bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], + payout_event_manager: PayoutEventManager, + adj_to_fail_with_tx_factory: Callable[..., None], ): # Now let's load it up and actually test some things delete_ledger_db() @@ -1083,9 +1093,9 @@ class TestProductCache: ) # 2. Payout - 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: Product, + product=product, amount=USDCent(71), ext_ref_id=uuid4().hex, created=start + timedelta(days=1, minutes=1), @@ -1104,11 +1114,11 @@ class TestProductCache: pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) product.set_cache( - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, ds=mnt_filepath, client=client_no_amm, bp_pem=brokerage_product_payout_event_manager, - redis_config=thl_redis_config: RedisConfig, + redis_config=thl_redis_config, ) # Fetch from cache and assert the instance loaded from redis diff --git a/tests/models/thl/test_product_userwalletconfig.py b/tests/models/thl/test_product_userwalletconfig.py index 4f6a6cc..b348981 100644 --- a/tests/models/thl/test_product_userwalletconfig.py +++ b/tests/models/thl/test_product_userwalletconfig.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from itertools import groupby from random import shuffle as rshuffle @@ -7,7 +9,7 @@ from generalresearch.models.thl.product import ( from generalresearch.models.thl.wallet import PayoutType -def all_equal(iterable): +def all_equal(iterable: list[str]) -> bool: g = groupby(iterable) return next(g, True) and not next(g, False) @@ -35,13 +37,13 @@ class TestProductUserWalletConfig: # in the same order because they're the same assert isinstance(instance.model_dump_json(), str) res = [] - for idx in range(100): + for _ in range(100): res.append(instance.model_dump_json()) assert all_equal(res) def test_model_dump_payout_types(self): res = [] - for idx in range(100): + for _ in range(100): # Generate a random order of PayoutTypes each time payout_types = [e for e in PayoutType] diff --git a/tests/models/thl/test_soft_pair.py b/tests/models/thl/test_soft_pair.py index 588847e..3cf835e 100644 --- a/tests/models/thl/test_soft_pair.py +++ b/tests/models/thl/test_soft_pair.py @@ -1,12 +1,14 @@ +from __future__ import annotations + from generalresearch.models import Source +from generalresearch.models.dynata.survey import ( + ConditionValueType, + DynataCondition, +) from generalresearch.models.thl.soft_pair import SoftPairResult, SoftPairResultType def test_model(): - from generalresearch.models.dynata.survey import ( - ConditionValueType, - DynataCondition, - ) c1 = DynataCondition( question_id="1", value_type=ConditionValueType.LIST, values=["a", "b"] diff --git a/tests/models/thl/test_upkquestion.py b/tests/models/thl/test_upkquestion.py index 99d7871..719fcff 100644 --- a/tests/models/thl/test_upkquestion.py +++ b/tests/models/thl/test_upkquestion.py @@ -1,13 +1,30 @@ +from __future__ import annotations + import pytest from pydantic import ValidationError +from generalresearch.models.morning.question import ( + MorningQuestion, + MorningQuestionType, +) +from generalresearch.models.thl.profiling.upk_question import ( + PatternValidation, + UPKImportance, + UpkQuestion, + UpkQuestionChoice, + UpkQuestionConfigurationMC, + UpkQuestionConfigurationTE, + UpkQuestionSelectorMC, + UpkQuestionSelectorTE, + UpkQuestionType, + UpkQuestionValidation, + order_exclusive_options, +) + class TestUpkQuestion: def test_importance(self): - from generalresearch.models.thl.profiling.upk_question import ( - UPKImportance, - ) res = UPKImportance(task_score=1, task_count=None) assert isinstance(res, UPKImportance) @@ -20,9 +37,6 @@ class TestUpkQuestion: assert "Input should be greater than or equal to 0" in str(e.value) def test_pattern(self): - from generalresearch.models.thl.profiling.upk_question import ( - PatternValidation, - ) s = PatternValidation(message="hi", pattern="x") with pytest.raises(ValidationError) as e: @@ -30,13 +44,6 @@ class TestUpkQuestion: assert "Instance is frozen" in str(e.value) def test_mc(self): - from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestion, - UpkQuestionChoice, - UpkQuestionConfigurationMC, - UpkQuestionSelectorMC, - UpkQuestionType, - ) q = UpkQuestion( id="601377a0d4c74529afc6293a8e5c3b5e", @@ -126,14 +133,6 @@ class TestUpkQuestion: assert "Extra inputs are not permitted" in str(e.value) def test_te(self): - from generalresearch.models.thl.profiling.upk_question import ( - PatternValidation, - UpkQuestion, - UpkQuestionConfigurationTE, - UpkQuestionSelectorTE, - UpkQuestionType, - UpkQuestionValidation, - ) q = UpkQuestion( id="601377a0d4c74529afc6293a8e5c3b5e", @@ -152,9 +151,6 @@ class TestUpkQuestion: assert q.choices is None def test_deserialization(self): - from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestion, - ) q = UpkQuestion.model_validate( { @@ -195,16 +191,18 @@ class TestUpkQuestion: assert q == UpkQuestion.model_validate(q.model_dump(mode="json")) def test_from_morning(self): - from generalresearch.models.morning.question import ( - MorningQuestion, - MorningQuestionType, - ) q = MorningQuestion( - id="gender", country_iso="us", language_iso="eng", name="Gender", text="What is your gender?", type="s", options=[ - {"id": "1", "text": "yes", "order": 1}, - {"id": "2", "text": "no", "order": 2}, - ] + id="gender", + country_iso="us", + language_iso="eng", + name="Gender", + text="What is your gender?", + type="s", + options=[ + {"id": "1", "text": "yes", "order": 1}, + {"id": "2", "text": "no", "order": 2}, + ], ) q.to_upk_question() q = MorningQuestion( @@ -218,13 +216,6 @@ class TestUpkQuestion: q.to_upk_question() def test_order(self): - from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestion, - UpkQuestionChoice, - UpkQuestionSelectorMC, - UpkQuestionType, - order_exclusive_options, - ) q = UpkQuestion( country_iso="us", @@ -258,9 +249,6 @@ class TestUpkQuestion: class TestUpkQuestionValidateAnswer: def test_validate_answer_SA(self): - from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestion, - ) question = UpkQuestion.model_validate( { @@ -296,9 +284,6 @@ class TestUpkQuestionValidateAnswer: ) def test_validate_answer_MA(self): - from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestion, - ) question = UpkQuestion.model_validate( { @@ -368,9 +353,6 @@ class TestUpkQuestionValidateAnswer: ) def test_validate_answer_TE(self): - from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestion, - ) question = UpkQuestion.model_validate( { diff --git a/tests/models/thl/test_user.py b/tests/models/thl/test_user.py index 0b8634a..9c4b548 100644 --- a/tests/models/thl/test_user.py +++ b/tests/models/thl/test_user.py @@ -1,4 +1,7 @@ +from __future__ import annotations + import json +from collections.abc import Callable from datetime import UTC, datetime, timedelta, timezone from decimal import Decimal from random import choice as rand_choice @@ -8,18 +11,21 @@ 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 + class TestUserUserID: def test_valid(self): - from generalresearch.models.thl.user import User val = randint(1, 2**30) user = User(user_id=val) assert user.user_id == val def test_type(self): - from generalresearch.models.thl.user import User # It will cast str to int assert User(user_id="1").user_id == 1 @@ -44,7 +50,6 @@ class TestUserUserID: assert "Input should be a valid integer," in str(cm.value) def test_zero(self): - from generalresearch.models.thl.user import User with pytest.raises(expected_exception=ValidationError) as cm: User(user_id=0) @@ -52,7 +57,6 @@ class TestUserUserID: assert "Input should be greater than 0" in str(cm.value) def test_negative(self): - from generalresearch.models.thl.user import User with pytest.raises(expected_exception=ValidationError) as cm: User(user_id=-1) @@ -60,7 +64,6 @@ class TestUserUserID: assert "Input should be greater than 0" in str(cm.value) def test_too_big(self): - from generalresearch.models.thl.user import User val = 2**31 with pytest.raises(expected_exception=ValidationError) as cm: @@ -69,7 +72,6 @@ class TestUserUserID: assert "Input should be less than 2147483648" in str(cm.value) def test_identifiable(self): - from generalresearch.models.thl.user import User val = randint(1, 2**30) user = User(user_id=val) @@ -80,7 +82,6 @@ class TestUserProductID: user_id = randint(1, 2**30) def test_valid(self): - from generalresearch.models.thl.user import User product_id = uuid4().hex @@ -89,7 +90,6 @@ class TestUserProductID: assert user.product_id == product_id def test_type(self): - from generalresearch.models.thl.user import User with pytest.raises(expected_exception=ValueError) as cm: User(user_id=self.user_id, product_id=0) @@ -102,7 +102,6 @@ class TestUserProductID: assert "Input should be a valid string" in str(cm.value) def test_empty(self): - from generalresearch.models.thl.user import User with pytest.raises(expected_exception=ValueError) as cm: User(user_id=self.user_id, product_id="") @@ -110,7 +109,6 @@ class TestUserProductID: assert "String should have at least 32 characters" in str(cm.value) def test_invalid_len(self): - from generalresearch.models.thl.user import User # Valid uuid4s are 32 char long product_id = uuid4().hex[:31] @@ -133,7 +131,6 @@ class TestUserProductID: assert "String should have at most 32 characters" in str(cm.value) def test_invalid_uuid(self): - from generalresearch.models.thl.user import User # Modify the UUID to break it product_id = uuid4().hex[:31] + "x" @@ -144,7 +141,6 @@ class TestUserProductID: assert "Invalid UUID" in str(cm.value) def test_invalid_hex_form(self): - from generalresearch.models.thl.user import User # Sure not in hex form, but it'll get caught for being the # wrong length before anything else @@ -157,7 +153,6 @@ class TestUserProductID: def test_identifiable(self): """Can't create a User with only a product_id because it also needs to the product_user_id""" - from generalresearch.models.thl.user import User product_id = uuid4().hex with pytest.raises(expected_exception=ValueError) as cm: @@ -172,10 +167,9 @@ class TestUserProductUserID: def randomword(self, length: int = 50): # Raw so nothing is escaped to add additional backslashes _bpuid_allowed = r"0123456789abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ!#$%&()*+,-.:;<=>?@[]^_{|}~" - return "".join(rand_choice(_bpuid_allowed) for i in range(length)) + return "".join(rand_choice(_bpuid_allowed) for _ in range(length)) def test_valid(self): - from generalresearch.models.thl.user import User product_user_id = uuid4().hex[:12] user = User(user_id=self.user_id, product_user_id=product_user_id) @@ -184,7 +178,6 @@ class TestUserProductUserID: assert user.product_user_id == product_user_id def test_type(self): - from generalresearch.models.thl.user import User with pytest.raises(expected_exception=ValueError) as cm: User(user_id=self.user_id, product_user_id=0) @@ -202,7 +195,6 @@ class TestUserProductUserID: assert "Input should be a valid string" in str(cm.value) def test_empty(self): - from generalresearch.models.thl.user import User with pytest.raises(expected_exception=ValueError) as cm: User(user_id=self.user_id, product_user_id="") @@ -210,7 +202,6 @@ class TestUserProductUserID: assert "String should have at least 3 characters" in str(cm.value) def test_invalid_len(self): - from generalresearch.models.thl.user import User product_user_id = self.randomword(251) with pytest.raises(expected_exception=ValueError) as cm: @@ -225,7 +216,6 @@ class TestUserProductUserID: assert "String should have at least 3 characters" in str(cm.value) def test_invalid_chars_space(self): - from generalresearch.models.thl.user import User product_user_id = f"{self.randomword(50)} {self.randomword(50)}" with pytest.raises(expected_exception=ValueError) as cm: @@ -234,7 +224,6 @@ class TestUserProductUserID: assert "String cannot contain spaces" in str(cm.value) def test_invalid_chars_slash(self): - from generalresearch.models.thl.user import User product_user_id = rf"{self.randomword(50)}\{self.randomword(50)}" with pytest.raises(expected_exception=ValueError) as cm: @@ -253,7 +242,6 @@ class TestUserProductUserID: I wanted a test that made sure the regex was hit. I do not know how we want to provide with the level of specific String checks we do in here for specific error messages.""" - from generalresearch.models.thl.user import User product_user_id = f"{self.randomword(50)}`{self.randomword(50)}" with pytest.raises(expected_exception=ValueError) as cm: @@ -275,7 +263,6 @@ class TestUserProductUserID: def test_identifiable(self): """Can't create a User with only a product_user_id because it also needs to the product_id""" - from generalresearch.models.thl.user import User product_user_id = uuid4().hex with pytest.raises(ValueError) as cm: @@ -288,7 +275,6 @@ class TestUserUUID: user_id = randint(1, 2**30) def test_valid(self): - from generalresearch.models.thl.user import User uuid_pk = uuid4().hex @@ -297,7 +283,6 @@ class TestUserUUID: assert user.uuid == uuid_pk def test_type(self): - from generalresearch.models.thl.user import User with pytest.raises(ValueError) as cm: User(user_id=self.user_id, uuid=0) @@ -315,7 +300,6 @@ class TestUserUUID: assert "Input should be a valid string" in str(cm.value) def test_empty(self): - from generalresearch.models.thl.user import User with pytest.raises(ValueError) as cm: User(user_id=self.user_id, uuid="") @@ -323,7 +307,6 @@ class TestUserUUID: assert "String should have at least 32 characters" in str(cm.value) def test_invalid_len(self): - from generalresearch.models.thl.user import User # Valid uuid4s are 32 char long uuid_pk = uuid4().hex[:31] @@ -341,7 +324,6 @@ class TestUserUUID: assert "String should have at most 32 characters" in str(cm.value) def test_invalid_uuid(self): - from generalresearch.models.thl.user import User # Modify the UUID to break it uuid_pk = uuid4().hex[:31] + "x" @@ -352,7 +334,6 @@ class TestUserUUID: assert "Invalid UUID" in str(cm.value) def test_invalid_hex_form(self): - from generalresearch.models.thl.user import User # Sure not in hex form, but it'll get caught for being the # wrong length before anything else @@ -369,7 +350,6 @@ class TestUserUUID: assert "Invalid UUID" in str(cm.value) def test_identifiable(self): - from generalresearch.models.thl.user import User user_uuid = uuid4().hex user = User(uuid=user_uuid) @@ -380,7 +360,6 @@ class TestUserCreated: user_id = randint(1, 2**30) def test_valid(self): - from generalresearch.models.thl.user import User user = User(user_id=self.user_id) dt = datetime.now(tz=UTC) @@ -389,7 +368,6 @@ class TestUserCreated: assert user.created == dt def test_tz_naive_throws_init(self): - from generalresearch.models.thl.user import User with pytest.raises(ValueError) as cm: User(user_id=self.user_id, created=datetime.now(tz=None)) # noqa @@ -397,7 +375,6 @@ class TestUserCreated: assert "Input should have timezone info" in str(cm.value) def test_tz_naive_throws_setter(self): - from generalresearch.models.thl.user import User user = User(user_id=self.user_id) with pytest.raises(ValueError) as cm: @@ -406,7 +383,6 @@ class TestUserCreated: assert "Input should have timezone info" in str(cm.value) def test_tz_utc(self): - from generalresearch.models.thl.user import User with pytest.raises(ValueError) as cm: User( @@ -417,7 +393,6 @@ class TestUserCreated: assert "Timezone is not UTC" in str(cm.value) def test_not_in_future(self): - from generalresearch.models.thl.user import User the_future = datetime.now(tz=UTC) + timedelta(minutes=1) with pytest.raises(ValueError) as cm: @@ -426,7 +401,6 @@ class TestUserCreated: assert "Input is in the future" in str(cm.value) def test_after_anno_domini(self): - from generalresearch.models.thl.user import User before_ad = datetime(year=2015, month=1, day=1, tzinfo=UTC) + timedelta( minutes=1 @@ -441,7 +415,6 @@ class TestUserLastSeen: user_id = randint(1, 2**30) def test_valid(self): - from generalresearch.models.thl.user import User user = User(user_id=self.user_id) dt = datetime.now(tz=UTC) @@ -450,7 +423,6 @@ class TestUserLastSeen: assert user.last_seen == dt def test_tz_naive_throws_init(self): - from generalresearch.models.thl.user import User with pytest.raises(ValueError) as cm: User(user_id=self.user_id, last_seen=datetime.now(tz=None)) # noqa @@ -458,7 +430,6 @@ class TestUserLastSeen: assert "Input should have timezone info" in str(cm.value) def test_tz_naive_throws_setter(self): - from generalresearch.models.thl.user import User user = User(user_id=self.user_id) with pytest.raises(ValueError) as cm: @@ -467,7 +438,6 @@ class TestUserLastSeen: assert "Input should have timezone info" in str(cm.value) def test_tz_utc(self): - from generalresearch.models.thl.user import User with pytest.raises(ValueError) as cm: User( @@ -478,7 +448,6 @@ class TestUserLastSeen: assert "Timezone is not UTC" in str(cm.value) def test_not_in_future(self): - from generalresearch.models.thl.user import User the_future = datetime.now(tz=UTC) + timedelta(minutes=1) with pytest.raises(ValueError) as cm: @@ -487,7 +456,6 @@ class TestUserLastSeen: assert "Input is in the future" in str(cm.value) def test_after_anno_domini(self): - from generalresearch.models.thl.user import User before_ad = datetime(year=2015, month=1, day=1, tzinfo=UTC) + timedelta( minutes=1 @@ -502,7 +470,6 @@ class TestUserBlocked: user_id = randint(1, 2**30) def test_valid(self): - from generalresearch.models.thl.user import User user = User(user_id=self.user_id, blocked=True) assert user.blocked @@ -510,7 +477,6 @@ class TestUserBlocked: def test_str_casting(self): """We don't want any of these to work, and that's why we set strict=True on the column""" - from generalresearch.models.thl.user import User with pytest.raises(ValueError) as cm: User(user_id=self.user_id, blocked="true") @@ -547,7 +513,6 @@ class TestUserTiming: user_id = randint(1, 2**30) def test_valid(self): - from generalresearch.models.thl.user import User created = datetime.now(tz=UTC) - timedelta(minutes=60) last_seen = datetime.now(tz=UTC) - timedelta(minutes=59) @@ -557,7 +522,6 @@ class TestUserTiming: assert user.last_seen == last_seen def test_created_first(self): - from generalresearch.models.thl.user import User created = datetime.now(tz=UTC) - timedelta(minutes=60) last_seen = datetime.now(tz=UTC) - timedelta(minutes=59) @@ -572,7 +536,6 @@ class TestUserModelVerification: """Tests that may be dependent on more than 1 attribute""" def test_identifiable(self): - from generalresearch.models.thl.user import User product_id = uuid4().hex product_user_id = uuid4().hex @@ -580,7 +543,6 @@ class TestUserModelVerification: assert user.is_identifiable def test_valid_helper(self): - from generalresearch.models.thl.user import User user_bool = User.is_valid_ubp( product_id=uuid4().hex, product_user_id=uuid4().hex @@ -594,7 +556,6 @@ class TestUserModelVerification: class TestUserSerialization: def test_basic_json(self): - from generalresearch.models.thl.user import User product_id = uuid4().hex product_user_id = uuid4().hex @@ -615,7 +576,6 @@ class TestUserSerialization: assert d.get("created").endswith("Z") def test_basic_dict(self): - from generalresearch.models.thl.user import User product_id = uuid4().hex product_user_id = uuid4().hex @@ -633,10 +593,11 @@ class TestUserSerialization: assert not d.get("blocked") assert d.get("product") is None - assert d.get("created").tzinfo == UTC + created = d.get("created") + assert isinstance(created, datetime) + assert created.tzinfo == UTC def test_from_json(self): - from generalresearch.models.thl.user import User product_id = uuid4().hex product_user_id = uuid4().hex @@ -651,12 +612,13 @@ class TestUserSerialization: u = User.model_validate_json(user.to_json()) assert u.product_id == product_id assert u.product is None + assert isinstance(u.created, datetime) assert u.created.tzinfo == UTC class TestUserMethods: - def test_audit_log(self, user, audit_log_manager): + def test_audit_log(self, user: User, audit_log_manager: AuditLogManager): assert user.audit_log is None user.prefetch_audit_log(audit_log_manager=audit_log_manager) assert user.audit_log == [] @@ -668,21 +630,21 @@ class TestUserMethods: def test_transactions( self, user_factory: Callable[..., User], - thl_lm, + thl_ledger_manager: ThlLedgerManager, session_with_tx_factory: Callable[..., None], - product_user_wallet_yes, + product_user_wallet_yes: Product, ): u1 = user_factory(product=product_user_wallet_yes) assert u1.transactions is None - u1.prefetch_transactions(thl_lm=thl_lm) + u1.prefetch_transactions(thl_lm=thl_ledger_manager) assert u1.transactions == [] session_with_tx_factory(user=u1) - u1.prefetch_transactions(thl_lm=thl_lm) + u1.prefetch_transactions(thl_lm=thl_ledger_manager) assert len(u1.transactions) == 1 @pytest.mark.skip(reason="TODO") - def test_location_history(self, user): + def test_location_history(self, user: User): assert user.location_history is None diff --git a/tests/models/thl/test_user_iphistory.py b/tests/models/thl/test_user_iphistory.py index d6ade9d..b8a0be3 100644 --- a/tests/models/thl/test_user_iphistory.py +++ b/tests/models/thl/test_user_iphistory.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from datetime import UTC, datetime, timedelta from generalresearch.models.thl.user_iphistory import ( diff --git a/tests/models/thl/test_user_metadata.py b/tests/models/thl/test_user_metadata.py index 3d851dc..a7b479d 100644 --- a/tests/models/thl/test_user_metadata.py +++ b/tests/models/thl/test_user_metadata.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import pytest from generalresearch.models import MAX_INT32 diff --git a/tests/models/thl/test_user_streak.py b/tests/models/thl/test_user_streak.py index 26c5e25..8300474 100644 --- a/tests/models/thl/test_user_streak.py +++ b/tests/models/thl/test_user_streak.py @@ -71,6 +71,7 @@ def test_user_streak_remaining(): ) print(f"{now.isoformat()=}, {end_of_today.isoformat()=}") expected = (end_of_today - now).total_seconds() + assert isinstance(us.time_remaining_in_period, timedelta) assert us.time_remaining_in_period.total_seconds() == pytest.approx(expected, abs=1) @@ -92,5 +93,6 @@ def test_user_streak_remaining_month(): ).replace(day=1) print(f"{now.isoformat()=}, {end_of_month.isoformat()=}") expected = (end_of_month - now).total_seconds() + assert isinstance(us.time_remaining_in_period, timedelta) assert us.time_remaining_in_period.total_seconds() == pytest.approx(expected, abs=1) print(us.time_remaining_in_period) diff --git a/tests/models/thl/test_wall.py b/tests/models/thl/test_wall.py index 88914ac..58e9825 100644 --- a/tests/models/thl/test_wall.py +++ b/tests/models/thl/test_wall.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from datetime import UTC, datetime, timedelta from decimal import Decimal from uuid import uuid4 diff --git a/tests/models/thl/test_wall_session.py b/tests/models/thl/test_wall_session.py index b39ad31..48b89ea 100644 --- a/tests/models/thl/test_wall_session.py +++ b/tests/models/thl/test_wall_session.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from datetime import UTC, datetime, timedelta from decimal import Decimal diff --git a/tests/wall_status_codes/test_analyze.py b/tests/wall_status_codes/test_analyze.py index fa53dbb..e36ca3d 100644 --- a/tests/wall_status_codes/test_analyze.py +++ b/tests/wall_status_codes/test_analyze.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from generalresearch.models.thl.definitions import Status, StatusCode1 from generalresearch.wall_status_codes import innovate diff --git a/tests/wxet/models/test_definitions.py b/tests/wxet/models/test_definitions.py index 543b9f1..a3616dd 100644 --- a/tests/wxet/models/test_definitions.py +++ b/tests/wxet/models/test_definitions.py @@ -1,12 +1,18 @@ +from __future__ import annotations + import pytest +from generalresearch.wxet.models.definitions import ( + WXETStatus, + WXETStatusCode1, + WXETStatusCode2, + check_wxet_status_consistent, +) + class TestWXETStatusCode1: def test_is_pre_task_entry_fail_pre(self): - from generalresearch.wxet.models.definitions import ( - WXETStatusCode1, - ) assert WXETStatusCode1.UNKNOWN.is_pre_task_entry_fail assert WXETStatusCode1.WXET_FAIL.is_pre_task_entry_fail @@ -32,12 +38,6 @@ class TestCheckWXETStatusConsistent: def test_completes(self): - from generalresearch.wxet.models.definitions import ( - WXETStatus, - WXETStatusCode1, - check_wxet_status_consistent, - ) - with pytest.raises(AssertionError) as cm: check_wxet_status_consistent( status=WXETStatus.COMPLETE, @@ -52,12 +52,6 @@ class TestCheckWXETStatusConsistent: def test_abandon(self): - from generalresearch.wxet.models.definitions import ( - WXETStatus, - WXETStatusCode1, - check_wxet_status_consistent, - ) - with pytest.raises(AssertionError) as cm: check_wxet_status_consistent( status=WXETStatus.ABANDON, @@ -71,12 +65,6 @@ class TestCheckWXETStatusConsistent: def test_fail(self): - from generalresearch.wxet.models.definitions import ( - WXETStatus, - WXETStatusCode1, - check_wxet_status_consistent, - ) - for sc1 in [ WXETStatusCode1.COMPLETE, WXETStatusCode1.WXET_ABANDON, @@ -95,13 +83,6 @@ class TestCheckWXETStatusConsistent: StatusCode1.WXET_FAIL """ - from generalresearch.wxet.models.definitions import ( - WXETStatus, - WXETStatusCode1, - WXETStatusCode2, - check_wxet_status_consistent, - ) - for sc2 in WXETStatusCode2: with pytest.raises(AssertionError) as cm: check_wxet_status_consistent( diff --git a/tests/wxet/models/test_finish_type.py b/tests/wxet/models/test_finish_type.py index 7bdeea7..afa3c76 100644 --- a/tests/wxet/models/test_finish_type.py +++ b/tests/wxet/models/test_finish_type.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import pytest from generalresearch.wxet.models.definitions import WXETStatus, WXETStatusCode1 -- cgit v1.2.3 From fdc170938ac4ac8fa5d9d4df1936e6dc0777f291 Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Thu, 27 Aug 2026 14:53:28 -0700 Subject: Ruff morning! Excluding thl_django in toml file --- generalresearch/models/thl/contest/__init__.py | 2 +- generalresearch/models/thl/task_status.py | 6 +- .../models/thl/wallet/cashout_method.py | 5 +- generalresearch/sql_helper.py | 6 +- generalresearch/wall_status_codes/dynata.py | 13 +- generalresearch/wall_status_codes/precision.py | 2 +- generalresearch/wall_status_codes/prodege.py | 2 +- generalresearch/wall_status_codes/repdata.py | 2 +- generalresearch/wall_status_codes/sago.py | 2 +- generalresearch/wall_status_codes/spectrum.py | 2 +- generalresearch/wall_status_codes/wxet.py | 4 +- test_utils/conftest.py | 19 ++- test_utils/models/contest/conftest.py | 24 ++-- test_utils/spectrum/conftest.py | 110 ++++++++++++++++ test_utils/spectrum/surveys_json.py | 140 --------------------- 15 files changed, 151 insertions(+), 188 deletions(-) delete mode 100644 test_utils/spectrum/surveys_json.py (limited to 'test_utils/models') diff --git a/generalresearch/models/thl/contest/__init__.py b/generalresearch/models/thl/contest/__init__.py index 0444586..842d85e 100644 --- a/generalresearch/models/thl/contest/__init__.py +++ b/generalresearch/models/thl/contest/__init__.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime from typing import Self from uuid import uuid4 diff --git a/generalresearch/models/thl/task_status.py b/generalresearch/models/thl/task_status.py index 011d743..de767d6 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, Self +from typing import Annotated, Any, Literal from pydantic import ( BaseModel, @@ -251,7 +251,7 @@ class TaskStatusResponse(BaseModel): return self.product_user_id @classmethod - def from_session(cls, session: Session, product: Product) -> Self: + def from_session(cls, session: Session, product: Product) -> TaskStatusResponse: user_payout_string = None if session.user_payout is not None: @@ -274,7 +274,7 @@ class TaskStatusResponse(BaseModel): user_payout_string=user_payout_string, product_id=session.user.product_id, product_user_id=session.user.product_user_id, - kwargs=session.url_metadata or dict(), + kwargs=session.url_metadata or {}, status_code_1=session.status_code_1, status_code_2=session.status_code_2, adjusted_status=session.adjusted_status, diff --git a/generalresearch/models/thl/wallet/cashout_method.py b/generalresearch/models/thl/wallet/cashout_method.py index 40c4717..6158afd 100644 --- a/generalresearch/models/thl/wallet/cashout_method.py +++ b/generalresearch/models/thl/wallet/cashout_method.py @@ -130,9 +130,8 @@ class CashoutMethodBase(BaseModel): f"Invalid amount requested: ${amount / 100:.2f}. Must be between" f" ${int(self.min_value) / 100:.2f} and ${int(self.max_value) / 100:.2f}" ) - if self.type == PayoutType.CASH_IN_MAIL: - if amount % 500 != 0: - raise ValueError("Amount must be in increments of $5.00") + if self.type == PayoutType.CASH_IN_MAIL and amount % 500 != 0: + raise ValueError("Amount must be in increments of $5.00") return True diff --git a/generalresearch/sql_helper.py b/generalresearch/sql_helper.py index ea45305..08b660d 100644 --- a/generalresearch/sql_helper.py +++ b/generalresearch/sql_helper.py @@ -166,7 +166,7 @@ class SqlHelper(SqlConnector): :param cursor: If cursor is passed, the insert is NOT committed! :param ignore_existing: adds 'ON CONFLICT DO NOTHING' to SQL statement. """ - assert len(set([len(x) for x in values_to_insert])) == 1 + assert len({len(x) for x in values_to_insert}) == 1 if cursor is None: connection = self.make_connection() c = connection.cursor() @@ -192,7 +192,6 @@ class SqlHelper(SqlConnector): if cursor is None: c.connection.commit() - def bulk_update( self, table_name: str, @@ -203,7 +202,7 @@ class SqlHelper(SqlConnector): if len(values_to_insert) == 0: return - assert len(set([len(x) for x in values_to_insert])) == 1 + assert len({len(x) for x in values_to_insert}) == 1 if cursor is None: connection = self.make_connection() c = connection.cursor() @@ -343,4 +342,3 @@ class SqlHelper(SqlConnector): c.execute(query) if cursor is None: c.connection.commit() - diff --git a/generalresearch/wall_status_codes/dynata.py b/generalresearch/wall_status_codes/dynata.py index 53eebe7..5d008c1 100644 --- a/generalresearch/wall_status_codes/dynata.py +++ b/generalresearch/wall_status_codes/dynata.py @@ -88,7 +88,7 @@ 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] = {} for k, v in status_codes_ext_map.items(): k: StatusCode1 v: list[str] @@ -117,16 +117,15 @@ def annotate_status_code( return status, status_code, None -def stop_marketplace_session(status_code_1: StatusCode1, ext_status_code_1) -> bool: +def stop_marketplace_session( + status_code_1: StatusCode1, ext_status_code_1: str +) -> bool: if ext_status_code_1.startswith("5"): # '5.10' is the user hit a Daily Limit, so they should not be sent in again today return True - if status_code_1 in { + return status_code_1 in { StatusCode1.PS_QUALITY, StatusCode1.BUYER_QUALITY_FAIL, StatusCode1.PS_BLOCKED, - }: - return True - - return False + } diff --git a/generalresearch/wall_status_codes/precision.py b/generalresearch/wall_status_codes/precision.py index c7593cf..4466fa6 100644 --- a/generalresearch/wall_status_codes/precision.py +++ b/generalresearch/wall_status_codes/precision.py @@ -73,7 +73,7 @@ status_codes_ext_map: dict[StatusCode1, list[str]] = { StatusCode1.PS_FAIL: ["21", "22"], StatusCode1.PS_OVERQUOTA: ["31", "32", "23"], } -ext_status_code_map = dict() +ext_status_code_map = {} for k, v in status_codes_ext_map.items(): k: StatusCode1 v: list[str] diff --git a/generalresearch/wall_status_codes/prodege.py b/generalresearch/wall_status_codes/prodege.py index 2876ba9..058ce28 100644 --- a/generalresearch/wall_status_codes/prodege.py +++ b/generalresearch/wall_status_codes/prodege.py @@ -31,7 +31,7 @@ status_code_map: dict[StatusCode1, list[str]] = { StatusCode1.PS_OVERQUOTA: ["13", "28", "29", "30", "31", "38"], } -status_class = dict() +status_class = {} for k, v in status_code_map.items(): k: StatusCode1 v: list[str] diff --git a/generalresearch/wall_status_codes/repdata.py b/generalresearch/wall_status_codes/repdata.py index 24aab14..0828b4c 100644 --- a/generalresearch/wall_status_codes/repdata.py +++ b/generalresearch/wall_status_codes/repdata.py @@ -58,7 +58,7 @@ status_code_map: dict[StatusCode1, list[str]] = { StatusCode1.PS_OVERQUOTA: ["5001", "6001", "6002", "6003"], } -status_class = dict() +status_class = {} for k, v in status_code_map.items(): k: StatusCode1 v: list[str] diff --git a/generalresearch/wall_status_codes/sago.py b/generalresearch/wall_status_codes/sago.py index b66190c..1a55a91 100644 --- a/generalresearch/wall_status_codes/sago.py +++ b/generalresearch/wall_status_codes/sago.py @@ -167,7 +167,7 @@ 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] = {} for k, v in status_codes_ext_map.items(): k: StatusCode1 v: list[str] diff --git a/generalresearch/wall_status_codes/spectrum.py b/generalresearch/wall_status_codes/spectrum.py index 610e239..7dee1e8 100644 --- a/generalresearch/wall_status_codes/spectrum.py +++ b/generalresearch/wall_status_codes/spectrum.py @@ -140,7 +140,7 @@ status_codes_ext_map: dict[StatusCode1, list[str]] = { "72", ], } -ext_status_code_map = dict() +ext_status_code_map = {} for k, v in status_codes_ext_map.items(): k: StatusCode1 v: list[str] diff --git a/generalresearch/wall_status_codes/wxet.py b/generalresearch/wall_status_codes/wxet.py index 6b3f67c..ad25ff3 100644 --- a/generalresearch/wall_status_codes/wxet.py +++ b/generalresearch/wall_status_codes/wxet.py @@ -29,7 +29,7 @@ status_codes_ext_map: dict[StatusCode1, list[WXETStatusCode1]] = { StatusCode1.UNKNOWN: [], StatusCode1.MARKETPLACE_FAIL: [WXETStatusCode1.BUYER_POSTBACK_NOT_RECEIVED], } -ext_status_code_map = dict() +ext_status_code_map = {} for k, v in status_codes_ext_map.items(): k: StatusCode1 v: list[WXETStatusCode1] @@ -58,7 +58,7 @@ status_code2_map: dict[StatusCode1, list[WXETStatusCode2]] = { WXETStatusCode2.TASK_VERSION_MISMATCH, ], } -ext_status_code2_map = dict() +ext_status_code2_map = {} for k, v in status_code2_map.items(): for vv in v: ext_status_code2_map[vv] = k diff --git a/test_utils/conftest.py b/test_utils/conftest.py index 03e9305..6894450 100644 --- a/test_utils/conftest.py +++ b/test_utils/conftest.py @@ -361,16 +361,15 @@ def delete_df_collection( create_main_accounts() case DFCollectionType.WALL | DFCollectionType.SESSION: - with thl_web_rw.make_connection() as conn: - with conn.cursor() as c: - c.execute("SET CONSTRAINTS ALL DEFERRED") - for table in [ - "thl_wall", - "thl_session", - ]: - c.execute( - query=f"DELETE FROM {table};", - ) + with thl_web_rw.make_connection() as conn, conn.cursor() as c: + c.execute("SET CONSTRAINTS ALL DEFERRED") + for table in [ + "thl_wall", + "thl_session", + ]: + c.execute( + query=f"DELETE FROM {table};", + ) case DFCollectionType.USER: for table in ["thl_usermetadata", "thl_user"]: diff --git a/test_utils/models/contest/conftest.py b/test_utils/models/contest/conftest.py index b14e126..d9c8a6b 100644 --- a/test_utils/models/contest/conftest.py +++ b/test_utils/models/contest/conftest.py @@ -11,7 +11,15 @@ 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, ) @@ -19,6 +27,8 @@ 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 @@ -33,18 +43,6 @@ from generalresearch.models.thl.user import User @pytest.fixture def raffle_contest_create() -> RaffleContestCreate: - from generalresearch.models.thl.contest import ( - ContestEndCondition, - ContestPrize, - ) - from generalresearch.models.thl.contest.definitions import ( - ContestPrizeKind, - ContestType, - ) - from generalresearch.models.thl.contest.raffle import ( - ContestEntryType, - RaffleContestCreate, - ) # This is what we'll get from the fastapi endpoint return RaffleContestCreate( @@ -89,7 +87,7 @@ def raffle_contest_factory( product_user_wallet_yes: Product, raffle_contest_create: RaffleContestCreate, contest_manager: ContestManager, -) -> Callable[..., Contest]: +) -> Callable[..., RaffleContest]: def _inner(**kwargs): raffle_contest_create.update(**kwargs) diff --git a/test_utils/spectrum/conftest.py b/test_utils/spectrum/conftest.py index d186a5b..7cd9321 100644 --- a/test_utils/spectrum/conftest.py +++ b/test_utils/spectrum/conftest.py @@ -88,6 +88,116 @@ def setup_spectrum_surveys( time.sleep(1) +@pytest.fixture(scope="session") +def spectrum_api_surveys_json() -> list[str]: + return [ + ( + '{"cpi":"3.90","country_isos":["us"],"language_isos":["eng"],"buyer_id":"215","bid_loi":780,"source":"s",' + '"used_question_ids":["1235","212"],"survey_id":"111111","survey_name":"Exciting New Survey #14472374",' + '"status":22,"field_end_date":"2023-03-02T07:05:36.261000Z","category_code":"232","calculation_type":"COMPLETES",' + '"requires_pii":false,"survey_exclusions":"13947261,14126487,14361592,14376811,14385771,14387789,14472374",' + '"exclusion_period":30,"bid_ir":0.2,"overall_loi":null,"overall_ir":null,"last_block_loi":null,' + '"last_block_ir":null,"project_last_complete_date":null,"country_iso":"us","language_iso":"eng",' + '"include_psids":null,"exclude_psids":null' + ',"qualifications":["ee5e842","e6e0b0b"],"quotas":[{"remaining_count":100,' + '"condition_hashes":["32cbf31"]}],"conditions":null,"created_api":"2023-02-28T07:05:36.698000Z",' + '"modified_api":"2024-03-10T09:43:40.030000Z","updated":"2024-05-30T21:52:46.431612Z","is_live":true' + "}" + ), + ( + '{"cpi":"3.90","country_isos":["us"],"language_isos":["eng"],"buyer_id":"215","bid_loi":780,"source":"s",' + '"used_question_ids":["1235","212"],"survey_id":"14472374","survey_name":"Exciting New Survey #14472374",' + '"status":22,"field_end_date":"2023-03-02T07:05:36.261000Z","category_code":"232","calculation_type":"COMPLETES",' + '"requires_pii":false,"survey_exclusions":"13947261,14126487,14361592,14376811,14385771,14387789,14472374",' + '"exclusion_period":30,"bid_ir":0.2,"overall_loi":null,"overall_ir":null,"last_block_loi":null,' + '"last_block_ir":null,"project_last_complete_date":null,"country_iso":"us","language_iso":"eng",' + '"include_psids":null,"exclude_psids":"0408319875e9dbffdc09e86671ad5636,23c4c66ecbc465906d0b0fd798740e64,' + '861df4603df3b7f754b8d4b89cbdb313","qualifications":["ee5e842","e6e0b0b"],"quotas":[{"remaining_count":100,' + '"condition_hashes":["32cbf31"]}],"conditions":null,"created_api":"2023-02-28T07:05:36.698000Z",' + '"modified_api":"2024-03-10T09:43:40.030000Z","updated":"2024-05-30T21:52:46.431612Z","is_live":true' + "}" + ), + ( + '{"cpi":"3.90","country_isos":["us"],"language_isos":["eng"],"buyer_id":"215","bid_loi":780,"source":"s",' + '"used_question_ids":["1235","212"],"survey_id":"12345","survey_name":"Exciting New Survey #14472374",' + '"status":22,"field_end_date":"2023-03-02T07:05:36.261000Z","category_code":"232","calculation_type":"COMPLETES",' + '"requires_pii":false,"survey_exclusions":"13947261,14126487,14361592,14376811,14385771,14387789,14472374",' + '"exclusion_period":30,"bid_ir":0.2,"overall_loi":null,"overall_ir":null,"last_block_loi":null,' + '"last_block_ir":null,"project_last_complete_date":null,"country_iso":"us","language_iso":"eng",' + '"include_psids":"7d043991b1494dbbb57786b11c88239c","exclude_psids":null' + ',"qualifications":["ee5e842","e6e0b0b"],"quotas":[{"remaining_count":100,' + '"condition_hashes":["32cbf31"]}],"conditions":null,"created_api":"2023-02-28T07:05:36.698000Z",' + '"modified_api":"2024-03-10T09:43:40.030000Z","updated":"2024-05-30T21:52:46.431612Z","is_live":true' + "}" + ), + ( + '{"cpi":"1.40","country_isos":["us"],"language_isos":["eng"],"buyer_id":"233","bid_loi":null,"source":"s",' + '"used_question_ids":["245","244","212","211","225"],"survey_id":"14970164","survey_name":"Exciting New Survey ' + '#14970164","status":22,"field_end_date":"2024-05-07T16:18:33.000000Z","category_code":"232",' + '"calculation_type":"COMPLETES","requires_pii":false,"survey_exclusions":"14970164,29690277",' + '"exclusion_period":30,"bid_ir":null,"overall_loi":900,"overall_ir":0.56,"last_block_loi":600,' + '"last_block_ir":0.01,"project_last_complete_date":"2024-05-28T04:12:56.297000Z","country_iso":"us",' + '"language_iso":"eng","include_psids":null,"exclude_psids":"01c7156fd9639737effbbdebd7fd66f6,' + "0508b88f4991bac8b10e9de74ce80194,0a51c627d77cef41f802e51a00126697,15b888176ac4781c2c978a9a05c396f8," + "17bc146b4f7fb05c7058d25da70c6a44,29935289c1f86a4144aab2e12652f305,2fe9d1d451efca10eba4fa4e5e2b74c9," + "c3527b7ef570a1571ea19870f3c25600,cdf2771d57cda9f1bf334382b2b7afd8,cebf3ec50395d973310ea526457dd5a0," + "cf3877cfc15e2e6ef2a56a7a7a37f3d3,dfa691e6d060e3643d5731df30be9f69,e0cb49537182660826aa351e1187809f," + 'edb6d280113ca49561f25fdcb500fde6,fbfba66cfad602f1c26e61e6174eb1f7,fd4307b16fd15e8534a4551c9b6872fc",' + '"qualifications":["1ab337d","a01aa68","437774f","dc6065b","82b6ad6"],"quotas":[{"remaining_count":242,' + '"condition_hashes":["c23c0b9"]},{"remaining_count":0,"condition_hashes":["5b8c6cf"]},{"remaining_count":126,' + '"condition_hashes":["ac35a6e"]},{"remaining_count":110,"condition_hashes":["5e7e5aa"]},{"remaining_count":108,' + '"condition_hashes":["9a7aef3"]},{"remaining_count":127,"condition_hashes":["4f75127"]},{"remaining_count":0,' + '"condition_hashes":["95437ed"]},{"remaining_count":17,"condition_hashes":["b4b7b95"]},{"remaining_count":16,' + '"condition_hashes":["0ab0ae6"]},{"remaining_count":8,"condition_hashes":["6e86fb5"]},{"remaining_count":12,' + '"condition_hashes":["24de31e"]},{"remaining_count":69,"condition_hashes":["6bdf350"]},{"remaining_count":411,' + '"condition_hashes":["c94d422"]}],"conditions":null,"created_api":"2023-03-30T22:47:36.324000Z",' + '"modified_api":"2024-05-30T13:07:16.489000Z","updated":"2024-05-30T21:52:37.493282Z","is_live":true,' + '"all_hashes":["c94d422","b4b7b95","6bdf350","6e86fb5","82b6ad6","24de31e","1ab337d","c23c0b9","9a7aef3",' + '"ac35a6e","95437ed","5b8c6cf","437774f","a01aa68","5e7e5aa","4f75127","0ab0ae6","dc6065b"]}' + ), + ( + '{"cpi":"1.23","country_isos":["au"],"language_isos":["eng"],"buyer_id":"215","bid_loi":780,"source":"s",' + '"used_question_ids":[],"survey_id":"69420","survey_name":"Everyone is eligible AU",' + '"status":22,"field_end_date":"2023-03-02T07:05:36.261000Z","category_code":"232","calculation_type":"COMPLETES",' + '"requires_pii":false,"survey_exclusions":"13947261,14126487,14361592,14376811,14385771,14387789,14472374",' + '"exclusion_period":30,"bid_ir":0.2,"overall_loi":null,"overall_ir":null,"last_block_loi":null,' + '"last_block_ir":null,"project_last_complete_date":null,"country_iso":"au","language_iso":"eng",' + '"include_psids":null,"exclude_psids":null' + ',"qualifications":[],"quotas":[{"remaining_count":100,' + '"condition_hashes":[]}],"conditions":null,"created_api":"2023-02-28T07:05:36.698000Z",' + '"modified_api":"2024-03-10T09:43:40.030000Z","updated":"2024-05-30T21:52:46.431612Z","is_live":true' + "}" + ), + ( + '{"cpi":"1.23","country_isos":["us"],"language_isos":["eng"],"buyer_id":"215","bid_loi":780,"source":"s",' + '"used_question_ids":[],"survey_id":"69421","survey_name":"Everyone is eligible US",' + '"status":22,"field_end_date":"2023-03-02T07:05:36.261000Z","category_code":"232","calculation_type":"COMPLETES",' + '"requires_pii":false,"survey_exclusions":"13947261,14126487,14361592,14376811,14385771,14387789,14472374",' + '"exclusion_period":30,"bid_ir":0.2,"overall_loi":null,"overall_ir":null,"last_block_loi":null,' + '"last_block_ir":null,"project_last_complete_date":null,"country_iso":"us","language_iso":"eng",' + '"include_psids":null,"exclude_psids":null' + ',"qualifications":[],"quotas":[{"remaining_count":100,' + '"condition_hashes":[]}],"conditions":null,"created_api":"2023-02-28T07:05:36.698000Z",' + '"modified_api":"2024-03-10T09:43:40.030000Z","updated":"2024-05-30T21:52:46.431612Z","is_live":true' + "}" + ), + # For partial eligibility + ( + '{"cpi":"1.23","country_isos":["us"],"language_isos":["eng"],"buyer_id":"215","bid_loi":780,"source":"s",' + '"used_question_ids":["1031", "212"],"survey_id":"999000","survey_name":"Pet owners",' + '"status":22,"field_end_date":"2023-03-02T07:05:36.261000Z","category_code":"232","calculation_type":"COMPLETES",' + '"requires_pii":false,"survey_exclusions":"13947261",' + '"exclusion_period":30,"bid_ir":0.2,"overall_loi":null,"overall_ir":null,"last_block_loi":null,' + '"last_block_ir":null,"project_last_complete_date":null,"country_iso":"us","language_iso":"eng",' + '"include_psids":null,"exclude_psids":null' + ',"qualifications":["0039b0c", "00f60a8"],"quotas":[{"remaining_count":100,' + '"condition_hashes":[]}],"conditions":null,"created_api":"2023-02-28T07:05:36.698000Z",' + '"modified_api":"2024-03-10T09:43:40.030000Z","updated":"2024-05-30T21:52:46.431612Z","is_live":true' + "}" + ), + ] + + @pytest.fixture(scope="session") def spectrum_api_survey_json() -> dict[str, Any]: return { diff --git a/test_utils/spectrum/surveys_json.py b/test_utils/spectrum/surveys_json.py deleted file mode 100644 index eb747a5..0000000 --- a/test_utils/spectrum/surveys_json.py +++ /dev/null @@ -1,140 +0,0 @@ -from generalresearch.models import LogicalOperator -from generalresearch.models.spectrum.survey import ( - SpectrumCondition, - SpectrumSurvey, -) -from generalresearch.models.thl.survey.condition import ConditionValueType - -SURVEYS_JSON = [ - '{"cpi":"3.90","country_isos":["us"],"language_isos":["eng"],"buyer_id":"215","bid_loi":780,"source":"s",' - '"used_question_ids":["1235","212"],"survey_id":"111111","survey_name":"Exciting New Survey #14472374",' - '"status":22,"field_end_date":"2023-03-02T07:05:36.261000Z","category_code":"232","calculation_type":"COMPLETES",' - '"requires_pii":false,"survey_exclusions":"13947261,14126487,14361592,14376811,14385771,14387789,14472374",' - '"exclusion_period":30,"bid_ir":0.2,"overall_loi":null,"overall_ir":null,"last_block_loi":null,' - '"last_block_ir":null,"project_last_complete_date":null,"country_iso":"us","language_iso":"eng",' - '"include_psids":null,"exclude_psids":null' - ',"qualifications":["ee5e842","e6e0b0b"],"quotas":[{"remaining_count":100,' - '"condition_hashes":["32cbf31"]}],"conditions":null,"created_api":"2023-02-28T07:05:36.698000Z",' - '"modified_api":"2024-03-10T09:43:40.030000Z","updated":"2024-05-30T21:52:46.431612Z","is_live":true' - "}", - '{"cpi":"3.90","country_isos":["us"],"language_isos":["eng"],"buyer_id":"215","bid_loi":780,"source":"s",' - '"used_question_ids":["1235","212"],"survey_id":"14472374","survey_name":"Exciting New Survey #14472374",' - '"status":22,"field_end_date":"2023-03-02T07:05:36.261000Z","category_code":"232","calculation_type":"COMPLETES",' - '"requires_pii":false,"survey_exclusions":"13947261,14126487,14361592,14376811,14385771,14387789,14472374",' - '"exclusion_period":30,"bid_ir":0.2,"overall_loi":null,"overall_ir":null,"last_block_loi":null,' - '"last_block_ir":null,"project_last_complete_date":null,"country_iso":"us","language_iso":"eng",' - '"include_psids":null,"exclude_psids":"0408319875e9dbffdc09e86671ad5636,23c4c66ecbc465906d0b0fd798740e64,' - '861df4603df3b7f754b8d4b89cbdb313","qualifications":["ee5e842","e6e0b0b"],"quotas":[{"remaining_count":100,' - '"condition_hashes":["32cbf31"]}],"conditions":null,"created_api":"2023-02-28T07:05:36.698000Z",' - '"modified_api":"2024-03-10T09:43:40.030000Z","updated":"2024-05-30T21:52:46.431612Z","is_live":true' - "}", - '{"cpi":"3.90","country_isos":["us"],"language_isos":["eng"],"buyer_id":"215","bid_loi":780,"source":"s",' - '"used_question_ids":["1235","212"],"survey_id":"12345","survey_name":"Exciting New Survey #14472374",' - '"status":22,"field_end_date":"2023-03-02T07:05:36.261000Z","category_code":"232","calculation_type":"COMPLETES",' - '"requires_pii":false,"survey_exclusions":"13947261,14126487,14361592,14376811,14385771,14387789,14472374",' - '"exclusion_period":30,"bid_ir":0.2,"overall_loi":null,"overall_ir":null,"last_block_loi":null,' - '"last_block_ir":null,"project_last_complete_date":null,"country_iso":"us","language_iso":"eng",' - '"include_psids":"7d043991b1494dbbb57786b11c88239c","exclude_psids":null' - ',"qualifications":["ee5e842","e6e0b0b"],"quotas":[{"remaining_count":100,' - '"condition_hashes":["32cbf31"]}],"conditions":null,"created_api":"2023-02-28T07:05:36.698000Z",' - '"modified_api":"2024-03-10T09:43:40.030000Z","updated":"2024-05-30T21:52:46.431612Z","is_live":true' - "}", - '{"cpi":"1.40","country_isos":["us"],"language_isos":["eng"],"buyer_id":"233","bid_loi":null,"source":"s",' - '"used_question_ids":["245","244","212","211","225"],"survey_id":"14970164","survey_name":"Exciting New Survey ' - '#14970164","status":22,"field_end_date":"2024-05-07T16:18:33.000000Z","category_code":"232",' - '"calculation_type":"COMPLETES","requires_pii":false,"survey_exclusions":"14970164,29690277",' - '"exclusion_period":30,"bid_ir":null,"overall_loi":900,"overall_ir":0.56,"last_block_loi":600,' - '"last_block_ir":0.01,"project_last_complete_date":"2024-05-28T04:12:56.297000Z","country_iso":"us",' - '"language_iso":"eng","include_psids":null,"exclude_psids":"01c7156fd9639737effbbdebd7fd66f6,' - "0508b88f4991bac8b10e9de74ce80194,0a51c627d77cef41f802e51a00126697,15b888176ac4781c2c978a9a05c396f8," - "17bc146b4f7fb05c7058d25da70c6a44,29935289c1f86a4144aab2e12652f305,2fe9d1d451efca10eba4fa4e5e2b74c9," - "c3527b7ef570a1571ea19870f3c25600,cdf2771d57cda9f1bf334382b2b7afd8,cebf3ec50395d973310ea526457dd5a0," - "cf3877cfc15e2e6ef2a56a7a7a37f3d3,dfa691e6d060e3643d5731df30be9f69,e0cb49537182660826aa351e1187809f," - 'edb6d280113ca49561f25fdcb500fde6,fbfba66cfad602f1c26e61e6174eb1f7,fd4307b16fd15e8534a4551c9b6872fc",' - '"qualifications":["1ab337d","a01aa68","437774f","dc6065b","82b6ad6"],"quotas":[{"remaining_count":242,' - '"condition_hashes":["c23c0b9"]},{"remaining_count":0,"condition_hashes":["5b8c6cf"]},{"remaining_count":126,' - '"condition_hashes":["ac35a6e"]},{"remaining_count":110,"condition_hashes":["5e7e5aa"]},{"remaining_count":108,' - '"condition_hashes":["9a7aef3"]},{"remaining_count":127,"condition_hashes":["4f75127"]},{"remaining_count":0,' - '"condition_hashes":["95437ed"]},{"remaining_count":17,"condition_hashes":["b4b7b95"]},{"remaining_count":16,' - '"condition_hashes":["0ab0ae6"]},{"remaining_count":8,"condition_hashes":["6e86fb5"]},{"remaining_count":12,' - '"condition_hashes":["24de31e"]},{"remaining_count":69,"condition_hashes":["6bdf350"]},{"remaining_count":411,' - '"condition_hashes":["c94d422"]}],"conditions":null,"created_api":"2023-03-30T22:47:36.324000Z",' - '"modified_api":"2024-05-30T13:07:16.489000Z","updated":"2024-05-30T21:52:37.493282Z","is_live":true,' - '"all_hashes":["c94d422","b4b7b95","6bdf350","6e86fb5","82b6ad6","24de31e","1ab337d","c23c0b9","9a7aef3",' - '"ac35a6e","95437ed","5b8c6cf","437774f","a01aa68","5e7e5aa","4f75127","0ab0ae6","dc6065b"]}', - '{"cpi":"1.23","country_isos":["au"],"language_isos":["eng"],"buyer_id":"215","bid_loi":780,"source":"s",' - '"used_question_ids":[],"survey_id":"69420","survey_name":"Everyone is eligible AU",' - '"status":22,"field_end_date":"2023-03-02T07:05:36.261000Z","category_code":"232","calculation_type":"COMPLETES",' - '"requires_pii":false,"survey_exclusions":"13947261,14126487,14361592,14376811,14385771,14387789,14472374",' - '"exclusion_period":30,"bid_ir":0.2,"overall_loi":null,"overall_ir":null,"last_block_loi":null,' - '"last_block_ir":null,"project_last_complete_date":null,"country_iso":"au","language_iso":"eng",' - '"include_psids":null,"exclude_psids":null' - ',"qualifications":[],"quotas":[{"remaining_count":100,' - '"condition_hashes":[]}],"conditions":null,"created_api":"2023-02-28T07:05:36.698000Z",' - '"modified_api":"2024-03-10T09:43:40.030000Z","updated":"2024-05-30T21:52:46.431612Z","is_live":true' - "}", - '{"cpi":"1.23","country_isos":["us"],"language_isos":["eng"],"buyer_id":"215","bid_loi":780,"source":"s",' - '"used_question_ids":[],"survey_id":"69421","survey_name":"Everyone is eligible US",' - '"status":22,"field_end_date":"2023-03-02T07:05:36.261000Z","category_code":"232","calculation_type":"COMPLETES",' - '"requires_pii":false,"survey_exclusions":"13947261,14126487,14361592,14376811,14385771,14387789,14472374",' - '"exclusion_period":30,"bid_ir":0.2,"overall_loi":null,"overall_ir":null,"last_block_loi":null,' - '"last_block_ir":null,"project_last_complete_date":null,"country_iso":"us","language_iso":"eng",' - '"include_psids":null,"exclude_psids":null' - ',"qualifications":[],"quotas":[{"remaining_count":100,' - '"condition_hashes":[]}],"conditions":null,"created_api":"2023-02-28T07:05:36.698000Z",' - '"modified_api":"2024-03-10T09:43:40.030000Z","updated":"2024-05-30T21:52:46.431612Z","is_live":true' - "}", - # For partial eligibility - '{"cpi":"1.23","country_isos":["us"],"language_isos":["eng"],"buyer_id":"215","bid_loi":780,"source":"s",' - '"used_question_ids":["1031", "212"],"survey_id":"999000","survey_name":"Pet owners",' - '"status":22,"field_end_date":"2023-03-02T07:05:36.261000Z","category_code":"232","calculation_type":"COMPLETES",' - '"requires_pii":false,"survey_exclusions":"13947261",' - '"exclusion_period":30,"bid_ir":0.2,"overall_loi":null,"overall_ir":null,"last_block_loi":null,' - '"last_block_ir":null,"project_last_complete_date":null,"country_iso":"us","language_iso":"eng",' - '"include_psids":null,"exclude_psids":null' - ',"qualifications":["0039b0c", "00f60a8"],"quotas":[{"remaining_count":100,' - '"condition_hashes":[]}],"conditions":null,"created_api":"2023-02-28T07:05:36.698000Z",' - '"modified_api":"2024-03-10T09:43:40.030000Z","updated":"2024-05-30T21:52:46.431612Z","is_live":true' - "}", -] - -# make sure hashes for 111111 are in db -c1 = SpectrumCondition( - question_id="1001", - value_type=ConditionValueType.LIST, - values=["a", "b", "c"], - negate=False, - logical_operator=LogicalOperator.OR, -) -c2 = SpectrumCondition( - question_id="1001", - value_type=ConditionValueType.LIST, - values=["a"], - negate=False, - logical_operator=LogicalOperator.OR, -) -c3 = SpectrumCondition( - question_id="1002", - value_type=ConditionValueType.RANGE, - values=["18-24", "30-32"], - negate=False, - logical_operator=LogicalOperator.OR, -) -c4 = SpectrumCondition( - question_id="212", - value_type=ConditionValueType.LIST, - values=["23", "24"], - negate=False, - logical_operator=LogicalOperator.OR, -) -c5 = SpectrumCondition( - question_id="1031", - value_type=ConditionValueType.LIST, - values=["113", "114", "121"], - negate=False, - logical_operator=LogicalOperator.OR, -) -CONDITIONS = [c1, c2, c3, c4, c5] -survey = SpectrumSurvey.model_validate_json(SURVEYS_JSON[0]) -assert c1.criterion_hash in survey.qualifications -assert c3.criterion_hash in survey.qualifications -- cgit v1.2.3 From ba544e2ba31432aad4d2acaba3e1f90c27137ded Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Fri, 28 Aug 2026 00:22:17 -0700 Subject: Ruff evening! --- generalresearch/managers/events.py | 40 +++---- generalresearch/managers/gr/authentication.py | 3 +- generalresearch/managers/innovate/survey.py | 2 +- generalresearch/managers/leaderboard/manager.py | 8 +- generalresearch/managers/morning/survey.py | 2 +- generalresearch/managers/network/label.py | 2 +- generalresearch/managers/network/mtr.py | 9 +- generalresearch/managers/network/nmap.py | 9 +- generalresearch/managers/network/rdns.py | 5 +- generalresearch/managers/sago/survey.py | 2 +- generalresearch/managers/spectrum/survey.py | 2 +- generalresearch/managers/thl/buyer.py | 4 +- generalresearch/managers/thl/cashout_method.py | 4 +- generalresearch/managers/thl/category.py | 7 +- generalresearch/managers/thl/contest_manager.py | 9 +- generalresearch/managers/thl/ipinfo.py | 8 +- .../managers/thl/ledger_manager/ledger.py | 13 ++- generalresearch/managers/thl/payout.py | 46 ++++---- generalresearch/managers/thl/product.py | 33 +++--- generalresearch/managers/thl/profiling/user_upk.py | 8 +- generalresearch/managers/thl/survey.py | 121 ++++++++++----------- generalresearch/managers/thl/survey_penalty.py | 1 - generalresearch/managers/thl/tango_api.py | 2 +- .../managers/thl/user_manager/__init__.py | 2 +- .../thl/user_manager/mysql_user_manager.py | 11 +- generalresearch/managers/thl/userhealth.py | 8 +- generalresearch/managers/thl/wall.py | 2 +- generalresearch/managers/thl/wallet/tango.py | 6 +- generalresearch/mariadb.py | 8 -- generalresearch/models/admin/__init__.py | 2 +- generalresearch/models/admin/request.py | 2 +- generalresearch/models/cint/question.py | 2 +- generalresearch/models/cint/survey.py | 15 +-- generalresearch/models/custom_types.py | 6 +- generalresearch/models/dynata/survey.py | 12 +- generalresearch/models/dynata/task_collection.py | 2 +- generalresearch/models/gr/authentication.py | 13 +-- generalresearch/models/gr/business.py | 64 +++++------ generalresearch/models/gr/team.py | 19 ++-- generalresearch/models/innovate/question.py | 10 +- generalresearch/models/innovate/survey.py | 39 ++++--- generalresearch/models/legacy/bucket.py | 30 +++-- generalresearch/models/legacy/questions.py | 32 ++---- generalresearch/models/morning/survey.py | 10 +- generalresearch/models/morning/task_collection.py | 2 +- generalresearch/models/network/label.py | 12 +- generalresearch/models/network/nmap/result.py | 2 +- generalresearch/models/network/rdns/command.py | 2 +- generalresearch/models/precision/question.py | 6 +- generalresearch/models/precision/survey.py | 2 +- generalresearch/models/prodege/question.py | 11 +- generalresearch/models/prodege/survey.py | 6 +- generalresearch/models/prodege/task_collection.py | 2 +- generalresearch/models/repdata/question.py | 4 +- generalresearch/models/repdata/survey.py | 7 +- generalresearch/models/repdata/task_collection.py | 2 +- generalresearch/models/sago/question.py | 5 +- generalresearch/models/sago/survey.py | 24 ++-- generalresearch/models/thl/contest/__init__.py | 2 +- generalresearch/models/thl/contest/contest.py | 2 +- .../models/thl/contest/contest_entry.py | 1 + generalresearch/models/thl/contest/raffle.py | 6 +- generalresearch/models/thl/demographics.py | 8 +- generalresearch/models/thl/finance.py | 2 +- generalresearch/models/thl/ledger.py | 4 +- generalresearch/models/thl/offerwall/__init__.py | 2 +- generalresearch/models/thl/offerwall/base.py | 6 +- generalresearch/models/thl/payout_format.py | 2 +- generalresearch/models/thl/product.py | 4 +- .../models/thl/profiling/marketplace.py | 7 +- generalresearch/models/thl/report_task.py | 2 +- generalresearch/models/thl/session.py | 11 +- generalresearch/models/thl/soft_pair.py | 2 +- generalresearch/models/thl/survey/penalty.py | 2 +- .../models/thl/survey/task_collection.py | 3 +- generalresearch/models/thl/task_status.py | 9 +- generalresearch/pg_helper.py | 6 +- generalresearch/sql_helper.py | 2 +- generalresearch/utils/grpc_logger.py | 6 +- generalresearch/wall_status_codes/fullcircle.py | 2 +- generalresearch/wall_status_codes/innovate.py | 2 +- generalresearch/wall_status_codes/lucid.py | 2 +- generalresearch/wall_status_codes/morning.py | 2 +- generalresearch/wall_status_codes/pollfish.py | 2 +- test_utils/managers/contest/conftest.py | 4 +- test_utils/models/contest/conftest.py | 37 ++++--- test_utils/spectrum/conftest.py | 77 ++++++++++--- .../incite/collections/test_df_collection_base.py | 17 ++- .../collections/test_df_collection_item_base.py | 19 +++- .../test_df_collection_thl_marketplaces.py | 14 ++- .../collections/test_df_collection_thl_web.py | 120 ++++++++++++++------ .../mergers/foundations/test_user_id_product.py | 21 +++- tests/incite/mergers/test_pop_ledger.py | 6 +- tests/incite/test_collection_base.py | 2 +- tests/managers/thl/test_contest/test_milestone.py | 11 +- .../test_ledger/test_thl_lm_tx__user_payouts.py | 22 ++-- tests/models/spectrum/test_survey.py | 44 +------- 97 files changed, 669 insertions(+), 546 deletions(-) (limited to 'test_utils/models') diff --git a/generalresearch/managers/events.py b/generalresearch/managers/events.py index f6c429e..efc8c0d 100644 --- a/generalresearch/managers/events.py +++ b/generalresearch/managers/events.py @@ -343,8 +343,8 @@ class TaskStatsManager(RedisManager): by_source=live_tasks_max_payout_by_source, ) - task_created_count_last_1h = dict() - task_created_count_last_24h = dict() + task_created_count_last_1h = {} + task_created_count_last_24h = {} for source in sources: task_created_count_last_1h[source] = pipe_res.pop(0) task_created_count_last_24h[source] = pipe_res.pop(0) @@ -381,25 +381,25 @@ class SessionStatsManager(RedisManager): older than 1 hr (in the 1 hr bucket) will expire. """ - # Must be ordered. Don't change this - global_keys = [ - "session_enters_last_1h", - "session_enters_last_24h", - "session_fails_last_1h", - "session_fails_last_24h", - "session_completes_last_1h", - "session_completes_last_24h", - "sum_payouts_last_1h", - "sum_payouts_last_24h", - "sum_user_payouts_last_1h", - "sum_user_payouts_last_24h", - # "session_fail_loi_sum_last_1h", - "session_fail_loi_sum_last_24h", - # "session_complete_loi_sum_last_1h", - "session_complete_loi_sum_last_24h", - ] - def __init__(self, *args, **kwargs): + # Must be ordered. Don't change this + self.global_keys = [ + "session_enters_last_1h", + "session_enters_last_24h", + "session_fails_last_1h", + "session_fails_last_24h", + "session_completes_last_1h", + "session_completes_last_24h", + "sum_payouts_last_1h", + "sum_payouts_last_24h", + "sum_user_payouts_last_1h", + "sum_user_payouts_last_24h", + # "session_fail_loi_sum_last_1h", + "session_fail_loi_sum_last_24h", + # "session_complete_loi_sum_last_1h", + "session_complete_loi_sum_last_24h", + ] + super().__init__(*args, **kwargs) self.SUM_HASH_LUA = self.redis_client.register_script(SUM_HASH_LUA_SCRIPT) diff --git a/generalresearch/managers/gr/authentication.py b/generalresearch/managers/gr/authentication.py index 80bee4b..721895e 100644 --- a/generalresearch/managers/gr/authentication.py +++ b/generalresearch/managers/gr/authentication.py @@ -270,7 +270,6 @@ class GRTokenManager(PostgresManager): ) conn.commit() - def get_by_user_id(self, user_id: PositiveInt) -> GRToken | None: # django authtoken_token table has (user_id) UNIQUE constraint # therefore, this will only return 0 or 1 GRTokens @@ -295,7 +294,7 @@ class GRTokenManager(PostgresManager): res = result[0] - for k, _ in res.items(): + for k in res: if isinstance(res[k], datetime): res[k] = res[k].replace(tzinfo=UTC) diff --git a/generalresearch/managers/innovate/survey.py b/generalresearch/managers/innovate/survey.py index cddfba2..f6d00a8 100644 --- a/generalresearch/managers/innovate/survey.py +++ b/generalresearch/managers/innovate/survey.py @@ -179,5 +179,5 @@ class InnovateSurveyManager(SurveyManager): if e.args[0] == 1062: existing_sns.add(sn) else: - raise e + raise self.update([surveys[sn] for sn in existing_sns]) diff --git a/generalresearch/managers/leaderboard/manager.py b/generalresearch/managers/leaderboard/manager.py index 07e3e2c..ed13cf2 100644 --- a/generalresearch/managers/leaderboard/manager.py +++ b/generalresearch/managers/leaderboard/manager.py @@ -45,9 +45,7 @@ class LeaderboardManager: self.country_iso = country_iso self.within_time_aware = None if within_time is None: - self.within_time_aware = datetime.now(tz=UTC).astimezone( - self.timezone - ) + self.within_time_aware = datetime.now(tz=UTC).astimezone(self.timezone) elif within_time.tzinfo is not None: self.within_time_aware = within_time.astimezone(self.timezone) else: @@ -123,7 +121,9 @@ class LeaderboardManager: user_idx = user_indices[0][0] user_row = user_indices[0][1] if user_row.rank == max([row.rank for row in rows]): - user_idx = [i for i, row in enumerate(rows) if row.rank == user_row.rank][0] + user_idx = next( + i for i, row in enumerate(rows) if row.rank == user_row.rank + ) start: int = max(user_idx - limit, 0) end: int = min(user_idx + limit + 1, len(rows)) diff --git a/generalresearch/managers/morning/survey.py b/generalresearch/managers/morning/survey.py index 0e29010..2488cd8 100644 --- a/generalresearch/managers/morning/survey.py +++ b/generalresearch/managers/morning/survey.py @@ -258,5 +258,5 @@ class MorningSurveyManager(SurveyManager): if e.args[0] == 1062: existing_sns.add(sn) else: - raise e + raise self.update([surveys[sn] for sn in existing_sns]) diff --git a/generalresearch/managers/network/label.py b/generalresearch/managers/network/label.py index 1f44862..f0ba9f7 100644 --- a/generalresearch/managers/network/label.py +++ b/generalresearch/managers/network/label.py @@ -30,7 +30,7 @@ class IPLabelManager(PostgresManager): params = ip_label.model_dump_postgres() with self.pg_config.make_connection() as conn, conn.cursor() as c: c.execute(query, params) - pk = c.fetchone()["id"] + _pk = c.fetchone()["id"] return ip_label def make_filter_str( diff --git a/generalresearch/managers/network/mtr.py b/generalresearch/managers/network/mtr.py index 19c5caf..54d74b7 100644 --- a/generalresearch/managers/network/mtr.py +++ b/generalresearch/managers/network/mtr.py @@ -42,8 +42,7 @@ class MTRRunManager(PostgresManager): if params_hops: c.executemany(query_hops, params_hops) else: - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute(query, params) - if params_hops: - c.executemany(query_hops, params_hops) + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute(query, params) + if params_hops: + c.executemany(query_hops, params_hops) diff --git a/generalresearch/managers/network/nmap.py b/generalresearch/managers/network/nmap.py index a8470c8..84d13ad 100644 --- a/generalresearch/managers/network/nmap.py +++ b/generalresearch/managers/network/nmap.py @@ -50,8 +50,7 @@ class NmapRunManager(PostgresManager): if nmap_run.ports: c.executemany(query_ports, params_ports) else: - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute(query, params) - if nmap_run.ports: - c.executemany(query_ports, params_ports) + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute(query, params) + if nmap_run.ports: + c.executemany(query_ports, params_ports) diff --git a/generalresearch/managers/network/rdns.py b/generalresearch/managers/network/rdns.py index 0b41a9a..c8ce913 100644 --- a/generalresearch/managers/network/rdns.py +++ b/generalresearch/managers/network/rdns.py @@ -28,6 +28,5 @@ class RDNSRunManager(PostgresManager): if c: c.execute(query, params) else: - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute(query, params) + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute(query, params) diff --git a/generalresearch/managers/sago/survey.py b/generalresearch/managers/sago/survey.py index 462d2ef..a13fbce 100644 --- a/generalresearch/managers/sago/survey.py +++ b/generalresearch/managers/sago/survey.py @@ -179,6 +179,6 @@ class SagoSurveyManager(SurveyManager): if e.args[0] == 1062: existing_sns.add(sn) else: - raise e + raise self.update([surveys[sn] for sn in existing_sns]) diff --git a/generalresearch/managers/spectrum/survey.py b/generalresearch/managers/spectrum/survey.py index 9b58d43..987ce7b 100644 --- a/generalresearch/managers/spectrum/survey.py +++ b/generalresearch/managers/spectrum/survey.py @@ -212,6 +212,6 @@ class SpectrumSurveyManager(SurveyManager): if e.args[0] == 1062: existing_sns.add(sn) else: - raise e + raise self.update([surveys[sn] for sn in existing_sns]) diff --git a/generalresearch/managers/thl/buyer.py b/generalresearch/managers/thl/buyer.py index 2cb582f..5aa2a01 100644 --- a/generalresearch/managers/thl/buyer.py +++ b/generalresearch/managers/thl/buyer.py @@ -18,8 +18,8 @@ class BuyerManager(PostgresManager): ): super().__init__(pg_config=pg_config, permissions=permissions) # self.buyer_pk: Dict[Buyer, int] = dict() - self.source_code_buyer: dict[str, Buyer] = dict() - self.source_code_pk: dict[str, int] = dict() + self.source_code_buyer: dict[str, Buyer] = {} + self.source_code_pk: dict[str, int] = {} self.populate_caches() def populate_caches(self): diff --git a/generalresearch/managers/thl/cashout_method.py b/generalresearch/managers/thl/cashout_method.py index 10282d4..f45e692 100644 --- a/generalresearch/managers/thl/cashout_method.py +++ b/generalresearch/managers/thl/cashout_method.py @@ -160,7 +160,7 @@ class CashoutMethodManager(PostgresManager): is_live: bool | None = True, ): filters = [] - params = dict() + params = {} if uuid is not None: params["uuid"] = uuid filters.append("id = %(uuid)s") @@ -292,7 +292,7 @@ class CashoutMethodManager(PostgresManager): x["type"] = PayoutType(x["provider"].upper()) if "data" not in x: - x["data"] = dict() + x["data"] = {} x["data"].update(x.pop("_data_")) x["data"]["type"] = x["type"] if user and x["type"] in {PayoutType.PAYPAL, PayoutType.CASH_IN_MAIL}: diff --git a/generalresearch/managers/thl/category.py b/generalresearch/managers/thl/category.py index 05ceb8f..e8a6aa6 100644 --- a/generalresearch/managers/thl/category.py +++ b/generalresearch/managers/thl/category.py @@ -9,8 +9,6 @@ from generalresearch.pg_helper import PostgresConfig class CategoryManager(PostgresManager): - categories = dict() - category_label_map = dict() def __init__( self, @@ -18,8 +16,9 @@ class CategoryManager(PostgresManager): permissions: Collection[Permission] | None = None, ): super().__init__(pg_config=pg_config, permissions=permissions) - self.categories: dict[UUIDStr, Category] = dict() - self.category_label_map: dict[str, Category] = dict() + self.categories: dict[UUIDStr, Category] = {} + self.category_label_map: dict[str, Category] = {} + self.populate_caches() def populate_caches(self): diff --git a/generalresearch/managers/thl/contest_manager.py b/generalresearch/managers/thl/contest_manager.py index 62146d7..64206e1 100644 --- a/generalresearch/managers/thl/contest_manager.py +++ b/generalresearch/managers/thl/contest_manager.py @@ -173,7 +173,7 @@ class ContestBaseManager(PostgresManager): except ValueError as e: if e.args[0] == "Contest not found": return None - raise e + raise @staticmethod def make_filter_str( @@ -187,7 +187,7 @@ class ContestBaseManager(PostgresManager): has_participants: bool | None = None, ) -> tuple[str, dict[str, Any]]: filters = [] - params = dict() + params = {} if product_id: params["product_id"] = product_id @@ -681,7 +681,7 @@ class RaffleContestManager(ContestBaseManager): raise ContestError(msg) if contest.entry_type == ContestEntryType.CASH: - tx = ledger_manager.create_tx_user_enter_contest( + ledger_manager.create_tx_user_enter_contest( contest_uuid=contest.uuid, contest_entry=entry ) @@ -827,7 +827,6 @@ class MilestoneContestManager(ContestBaseManager): ) self.end_milestone_contest(contest) - def enter_contest_db_work_milestone( self, contest: MilestoneUserView, user: User, incr: PositiveInt ) -> MilestoneEntry: @@ -1052,7 +1051,7 @@ class ContestManager( ) -> NonNegativeInt: contests_closed = 0 for contest in contests: - should_end, reason = contest.should_end() + should_end, _ = contest.should_end() if should_end: if hasattr(contest, "redis_client"): contest.redis_client = redis_client diff --git a/generalresearch/managers/thl/ipinfo.py b/generalresearch/managers/thl/ipinfo.py index e1143c2..9595cab 100644 --- a/generalresearch/managers/thl/ipinfo.py +++ b/generalresearch/managers/thl/ipinfo.py @@ -423,7 +423,9 @@ class IPInformationManager(PostgresManager): FROM thl_ipinformation WHERE updated >= NOW() - INTERVAL '12 hours' """ - denominator = list(pg_config.execute_sql_query(query=query))[0]["denominator"] + denominator = next(iter(pg_config.execute_sql_query(query=query)))[ + "denominator" + ] if denominator == 0: pass @@ -509,7 +511,7 @@ class GeoIpInfoManager(PostgresManagerWithRedis): res = [GeoIPInformation.model_validate_json(raw) for raw in res if raw] gs = {x.ip: x for x in res} - res2 = dict() + res2 = {} for ip, (normalized_ip, lookup_prefix) in ip_norm_lookup.items(): if normalized_ip not in gs: # try the non-normalized (remove me also 28 days from 2025-11-15) @@ -719,7 +721,7 @@ class GeoIpInfoManager(PostgresManagerWithRedis): gs = [GeoIPInformation.from_mysql(i) for i in res] gs = {g.ip: g for g in gs} - res2 = dict() + res2 = {} for ip, (normalized_ip, lookup_prefix) in ip_norm_lookup.items(): if normalized_ip not in gs: diff --git a/generalresearch/managers/thl/ledger_manager/ledger.py b/generalresearch/managers/thl/ledger_manager/ledger.py index 410f6ca..a5263f0 100644 --- a/generalresearch/managers/thl/ledger_manager/ledger.py +++ b/generalresearch/managers/thl/ledger_manager/ledger.py @@ -150,7 +150,7 @@ class LedgerTransactionManager(LedgerManagerBasePostgres): ), "LedgerTransactionManager has insufficient Permissions" if metadata is None: - metadata = dict() + metadata = {} if created is None: created = datetime.now(tz=UTC) @@ -429,7 +429,7 @@ class LedgerTransactionManager(LedgerManagerBasePostgres): ) } else: - metadata = dict() + metadata = {} entries = [ LedgerEntry( @@ -750,7 +750,7 @@ class LedgerMetadataManager(LedgerManagerBasePostgres): """ - tx_ids = set([tx.id for tx in transactions]) + tx_ids = {tx.id for tx in transactions} res = self.pg_config.execute_sql_query( query=""" SELECT @@ -782,7 +782,7 @@ class LedgerMetadataManager(LedgerManagerBasePostgres): from the database. """ - tx_ids = set([tx.id for tx in transactions]) + tx_ids = {tx.id for tx in transactions} res = self.pg_config.execute_sql_query( query=""" SELECT tx_meta.id @@ -792,7 +792,7 @@ class LedgerMetadataManager(LedgerManagerBasePostgres): params=[list(tx_ids)], ) - return set([i["id"] for i in res]) + return {i["id"] for i in res} class LedgerEntryManager(LedgerManagerBasePostgres): @@ -803,7 +803,7 @@ class LedgerEntryManager(LedgerManagerBasePostgres): def get_tx_entries_by_txs( self, transactions: list[LedgerTransaction] ) -> list[LedgerEntry]: - tx_ids = set([tx.id for tx in transactions]) + tx_ids = {tx.id for tx in transactions} tx_entries = self.pg_config.execute_sql_query( query=""" SELECT @@ -1141,4 +1141,5 @@ class LedgerManager( } for k, v in d.items(): v["total"] = (v["debit"] - v["credit"]) * k.normal_balance.value + return d diff --git a/generalresearch/managers/thl/payout.py b/generalresearch/managers/thl/payout.py index 8bc6843..f50e0d2 100644 --- a/generalresearch/managers/thl/payout.py +++ b/generalresearch/managers/thl/payout.py @@ -95,9 +95,9 @@ class PayoutEventManager(PostgresManagerWithRedis): with self.pg_config.make_connection() as conn: with conn.cursor() as c: c.execute(query=query, params=d) - assert c.rowcount == 1, ( - "Nothing was updated! Are you sure this payout_event exists?" - ) + assert ( + c.rowcount == 1 + ), "Nothing was updated! Are you sure this payout_event exists?" conn.commit() @@ -140,7 +140,7 @@ class UserPayoutEventManager(PayoutEventManager): # the purposes of returning to the user. pe = self.get_by_uuid(pe_uuid=pe_uuid) - transaction_info = dict() + transaction_info = {} order: dict[str, Any] = pe.order_data if pe.payout_type == PayoutType.TANGO and pe.status == PayoutStatus.COMPLETE: reward = order["reward"] @@ -411,7 +411,7 @@ class BrokerageProductPayoutEventManager(PayoutEventManager): *** IT IS ONLY FOR Brokerage Product PAYOUTS *** """ - params = dict() + params = {} filters = [] if ext_ref_id: # This is transaction id for tracking ACH/Wires with a banking institution @@ -630,9 +630,9 @@ class BusinessPayoutEventManager(PostgresManagerWithRedis): for bp_payout in d["bp_payouts"]: bp_payout["created"] = datetime.fromisoformat(bp_payout["created"]) bpe = BusinessPayoutEvent.model_validate(d) - assert bpe.bp_payouts is not None and len(bpe.bp_payouts) > 0, ( - "No BP payouts found for this Business Payout Event. This shouldn't happen!" - ) + assert ( + bpe.bp_payouts is not None and len(bpe.bp_payouts) > 0 + ), "No BP payouts found for this Business Payout Event. This shouldn't happen!" return bpe def filter_by( @@ -677,9 +677,9 @@ class BusinessPayoutEventManager(PostgresManagerWithRedis): for bp_payout in row["bp_payouts"]: bp_payout["created"] = datetime.fromisoformat(bp_payout["created"]) bpe = BusinessPayoutEvent.model_validate(row) - assert bpe.bp_payouts is not None and len(bpe.bp_payouts) > 0, ( - "No BP payouts found for this Business Payout Event. This shouldn't happen!" - ) + assert ( + bpe.bp_payouts is not None and len(bpe.bp_payouts) > 0 + ), "No BP payouts found for this Business Payout Event. This shouldn't happen!" bpes.append(bpe) return bpes @@ -696,9 +696,9 @@ class BusinessPayoutEventManager(PostgresManagerWithRedis): for bp_pe in bpe.bp_payouts ] txs = thl_lm.get_tx_ids_by_tags(tags=tags) - assert len(txs) == len(bpe.bp_payouts), ( - f"Expected {len(bpe.bp_payouts)} BP payouts but found {len(txs)}!" - ) + assert len(txs) == len( + bpe.bp_payouts + ), f"Expected {len(bpe.bp_payouts)} BP payouts but found {len(txs)}!" return True def resume_failed_business_payout( @@ -824,9 +824,9 @@ class BusinessPayoutEventManager(PostgresManagerWithRedis): shortfall: int = int(target_amount) - w_df["deduction"].sum() w_df["remaining_balance"] = w_df["available_balance"] - w_df["deduction"] - assert w_df[w_df["deduction"] > w_df["available_balance"]].empty, ( - "Trying to deduct more from an Product than what is available" - ) + assert w_df[ + w_df["deduction"] > w_df["available_balance"] + ].empty, "Trying to deduct more from an Product than what is available" return w_df @@ -898,9 +898,9 @@ class BusinessPayoutEventManager(PostgresManagerWithRedis): ) -> BusinessPayoutEvent: assert isinstance(bpe, BusinessPayoutEventCreate) assert bpe.bp_payouts, "Must provide at least one BP Payout" - assert {bp_pe.status for bp_pe in bpe.bp_payouts} == {PayoutStatus.PENDING}, ( - "All BP Payouts must be PENDING" - ) + assert {bp_pe.status for bp_pe in bpe.bp_payouts} == { + PayoutStatus.PENDING + }, "All BP Payouts must be PENDING" INSERT_SUPPLIER_PAYOUT = """ INSERT INTO supplier_payout ( business_id, created, amount, @@ -993,9 +993,9 @@ class BusinessPayoutEventManager(PostgresManagerWithRedis): # Can't pay any Products that don't have a remaining balance df = df[df["remaining_balance"] > 0].copy() - assert df.deduction.sum() == business.balance.recoup, ( - "recoup_proportional failure" - ) + assert ( + df.deduction.sum() == business.balance.recoup + ), "recoup_proportional failure" df["issue_amount"] = BusinessPayoutEventManager.distribute_amount( df=df, amount=amount diff --git a/generalresearch/managers/thl/product.py b/generalresearch/managers/thl/product.py index 54fa7c8..aac2979 100644 --- a/generalresearch/managers/thl/product.py +++ b/generalresearch/managers/thl/product.py @@ -33,6 +33,7 @@ if TYPE_CHECKING: ProfilingConfig, SessionConfig, SourcesConfig, + SupplyConfig, UserCreateConfig, UserHealthConfig, UserWalletConfig, @@ -167,15 +168,14 @@ class ProductManager(PostgresManager): if filter_uuids is None or len(filter_uuids) == 0: return [] - with self.pg_config.make_connection() as sql_connection: - with sql_connection.cursor() as c: - res = [] - for chunk in chunked(filter_uuids, 500): - res.extend( - self.fetch_uuids_( - c=c, filter_uuids=chunk, filter_column=filter_column - ) + with self.pg_config.make_connection() as sql_connection, sql_connection.cursor() as c: + res = [] + for chunk in chunked(filter_uuids, 500): + res.extend( + self.fetch_uuids_( + c=c, filter_uuids=chunk, filter_column=filter_column ) + ) return res def fetch_uuids_( @@ -258,9 +258,10 @@ class ProductManager(PostgresManager): for k, v in res1.items(): try: r.append(Product.model_validate(v)) - except ValidationError as e: + except ValidationError: logger.info(f"failed to parse product: {k}") - raise e + raise + return r def create( @@ -272,7 +273,7 @@ class ProductManager(PostgresManager): business_id: UUIDStr | None = None, harmonizer_domain: str | None = None, commission_pct: Decimal = Decimal("0.05"), - sources_config: SourcesConfig | SupplyConfigs | None = None, + sources_config: SourcesConfig | SupplyConfig | None = None, payout_config: PayoutConfig | None = None, session_config: SessionConfig | None = None, profiling_config: ProfilingConfig | None = None, @@ -360,10 +361,10 @@ class ProductManager(PostgresManager): insert_data["payments_enabled"] = instance.payments_enabled try: - insert_data["id_int"] = list(self.pg_config.execute_sql_query(query=""" + insert_data["id_int"] = next(iter(self.pg_config.execute_sql_query(query=""" SELECT COALESCE(MAX(id_int), 0) + 1 as id_int FROM userprofile_brokerageproduct - """))[0]["id_int"] + """)))["id_int"] instance.id_int = insert_data["id_int"] query = """ @@ -400,14 +401,14 @@ class ProductManager(PostgresManager): try: return self.get_by_uuid(product_uuid=instance.id) - except Exception: + except AssertionError: pass finally: self.cache_clear(instance.id) # If we couldn't find the Product, then go ahead and raise. capture_exception(e) - raise e + raise bpconfig = instance.model_dump( include={"sources_config", "user_wallet"}, mode="json" @@ -477,7 +478,7 @@ class ProductManager(PostgresManager): data["grs_domain"] = data.pop("harmonizer_domain") data = {k: v for k, v in data.items() if k in in_bp_keys} data["id"] = product_uuid - update_str = ", ".join(f"{k}=%({k})s" for k in data.keys()) + update_str = ", ".join(f"{k}=%({k})s" for k in data) self.pg_config.execute_write( f""" UPDATE userprofile_brokerageproduct diff --git a/generalresearch/managers/thl/profiling/user_upk.py b/generalresearch/managers/thl/profiling/user_upk.py index a2cddb3..6106037 100644 --- a/generalresearch/managers/thl/profiling/user_upk.py +++ b/generalresearch/managers/thl/profiling/user_upk.py @@ -158,7 +158,7 @@ class UserUpkManager(PostgresManagerWithRedis): country_isos = {x["country_iso"] for x in upk_ans_dict} assert len(country_isos) == 1 - country_iso = list(country_isos)[0] + country_iso = next(iter(country_isos)) for x in upk_ans_dict: x["pred"] = x["pred"].replace("gr:", "") x["obj"] = x["obj"].replace("gr:", "") @@ -304,15 +304,15 @@ class UserUpkManager(PostgresManagerWithRedis): def set_user_upk(self, upk_ans: list[UpkQuestionAnswer]) -> None: user_id = {x.user_id for x in upk_ans} assert len(user_id) == 1, "only run for 1 user at a time" - user_id = list(user_id)[0] + user_id = next(iter(user_id)) curr_upk = self.get_user_upk(user_id=user_id) curr_upk_simple = self.get_user_upk_simple(user_id=user_id) new_upk_simple = defaultdict(set) delete_items = set() - upk_multi = list() - delete_upk_multi = list() + upk_multi = [] + delete_upk_multi = [] for x in upk_ans: # For zero or more (multiple values) We want all values to equal these. # Might involve deleting values if they exist and are not in upk_ans diff --git a/generalresearch/managers/thl/survey.py b/generalresearch/managers/thl/survey.py index a9ec841..024ad38 100644 --- a/generalresearch/managers/thl/survey.py +++ b/generalresearch/managers/thl/survey.py @@ -134,7 +134,7 @@ class SurveyManager(PostgresManager): if len(survey_keys) == 0: return [] - params = dict() + params = {} survey_source_ids = defaultdict(set) for sk in survey_keys: @@ -354,59 +354,6 @@ class SurveyManager(PostgresManager): class SurveyStatManager(PostgresManager): - KEYS = [ - "survey_id", - "quota_id", - "country_iso", - "version", - "cpi", - "complete_too_fast_cutoff", - "prescreen_conv_alpha", - "prescreen_conv_beta", - "conv_alpha", - "conv_beta", - "dropoff_alpha", - "dropoff_beta", - "completion_time_mu", - "completion_time_sigma", - "mobile_eligible_alpha", - "mobile_eligible_beta", - "desktop_eligible_alpha", - "desktop_eligible_beta", - "tablet_eligible_alpha", - "tablet_eligible_beta", - "long_fail_rate", - "user_report_coeff", - "recon_likelihood", - "score_x0", - "score_x1", - "score", - "updated_at", - "survey_is_live", - "survey_survey_id", - "survey_source", - ] - - SURVEY_STATS_COL_MAP = { - "PRESCREEN_CONVERSION.alpha": "prescreen_conv_alpha", - "PRESCREEN_CONVERSION.beta": "prescreen_conv_beta", - "CONVERSION.alpha": "conv_alpha", - "CONVERSION.beta": "conv_beta", - "COMPLETION_TIME.mu": "completion_time_mu", - "COMPLETION_TIME.sigma": "completion_time_sigma", - "LONG_FAIL.value": "long_fail_rate", - "USER_REPORT_COEFF.value": "user_report_coeff", - "RECON_LIKELIHOOD.value": "recon_likelihood", - "DROPOFF_RATE.alpha": "dropoff_alpha", - "DROPOFF_RATE.beta": "dropoff_beta", - "IS_MOBILE_ELIGIBLE.alpha": "mobile_eligible_alpha", - "IS_MOBILE_ELIGIBLE.beta": "mobile_eligible_beta", - "IS_DESKTOP_ELIGIBLE.alpha": "desktop_eligible_alpha", - "IS_DESKTOP_ELIGIBLE.beta": "desktop_eligible_beta", - "IS_TABLET_ELIGIBLE.alpha": "tablet_eligible_alpha", - "IS_TABLET_ELIGIBLE.beta": "tablet_eligible_beta", - "cpi": "cpi", - } def __init__( self, @@ -419,6 +366,60 @@ class SurveyStatManager(PostgresManager): ) # self.ensure_surveystat_key_type() + self.KEYS = [ + "survey_id", + "quota_id", + "country_iso", + "version", + "cpi", + "complete_too_fast_cutoff", + "prescreen_conv_alpha", + "prescreen_conv_beta", + "conv_alpha", + "conv_beta", + "dropoff_alpha", + "dropoff_beta", + "completion_time_mu", + "completion_time_sigma", + "mobile_eligible_alpha", + "mobile_eligible_beta", + "desktop_eligible_alpha", + "desktop_eligible_beta", + "tablet_eligible_alpha", + "tablet_eligible_beta", + "long_fail_rate", + "user_report_coeff", + "recon_likelihood", + "score_x0", + "score_x1", + "score", + "updated_at", + "survey_is_live", + "survey_survey_id", + "survey_source", + ] + + self.SURVEY_STATS_COL_MAP = { + "PRESCREEN_CONVERSION.alpha": "prescreen_conv_alpha", + "PRESCREEN_CONVERSION.beta": "prescreen_conv_beta", + "CONVERSION.alpha": "conv_alpha", + "CONVERSION.beta": "conv_beta", + "COMPLETION_TIME.mu": "completion_time_mu", + "COMPLETION_TIME.sigma": "completion_time_sigma", + "LONG_FAIL.value": "long_fail_rate", + "USER_REPORT_COEFF.value": "user_report_coeff", + "RECON_LIKELIHOOD.value": "recon_likelihood", + "DROPOFF_RATE.alpha": "dropoff_alpha", + "DROPOFF_RATE.beta": "dropoff_beta", + "IS_MOBILE_ELIGIBLE.alpha": "mobile_eligible_alpha", + "IS_MOBILE_ELIGIBLE.beta": "mobile_eligible_beta", + "IS_DESKTOP_ELIGIBLE.alpha": "desktop_eligible_alpha", + "IS_DESKTOP_ELIGIBLE.beta": "desktop_eligible_beta", + "IS_TABLET_ELIGIBLE.alpha": "tablet_eligible_alpha", + "IS_TABLET_ELIGIBLE.beta": "tablet_eligible_beta", + "cpi": "cpi", + } + # # def ensure_surveystat_key_type(self): # SQL = """ @@ -570,12 +571,10 @@ class SurveyStatManager(PostgresManager): = (v.survey_id, v.quota_id, v.country_iso, v.version); """ params = [item for row in keys for item in row] - with self.pg_config.make_connection() as conn: - # self.register_surveystat_key(conn) - with conn.cursor() as c: - c.execute(query, params=params) - res = c.fetchall() - # print('\n'.join([x['QUERY PLAN'] for x in res])) + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute(query, params=params) + res = c.fetchall() + # print('\n'.join([x['QUERY PLAN'] for x in res])) return [SurveyStat.model_validate(x) for x in res] def update_surveystats_for_source( @@ -633,7 +632,7 @@ class SurveyStatManager(PostgresManager): country_iso: str | None = None, ) -> tuple[str, dict[str, Any]]: filters = [] - params = dict() + params = {} if updated_after is not None: params["updated_after"] = updated_after filters.append("ss.updated_at >= %(updated_after)s") diff --git a/generalresearch/managers/thl/survey_penalty.py b/generalresearch/managers/thl/survey_penalty.py index efaa930..bf914cb 100644 --- a/generalresearch/managers/thl/survey_penalty.py +++ b/generalresearch/managers/thl/survey_penalty.py @@ -61,7 +61,6 @@ class SurveyPenaltyManager(RedisManager): return f"{self.redis_prefix}:{uuid_id}" def set_penalties(self, penalties: list[Penalty]): - """ """ if len(penalties) > 1000: LOG.warning("SurveyPenaltyManager.set_penalties batch me!") assert len(penalties) < 10_000, "something is surely wrong" diff --git a/generalresearch/managers/thl/tango_api.py b/generalresearch/managers/thl/tango_api.py index 657224e..dab560e 100644 --- a/generalresearch/managers/thl/tango_api.py +++ b/generalresearch/managers/thl/tango_api.py @@ -122,7 +122,7 @@ class TangoClient: return self.get_order(reference_order_id) except TangoError as e: if "The order you requested cannot be found" not in e.args[0]: - raise e + raise return None def create_order(self, order: TangoOrderRequest) -> dict[str, Any]: diff --git a/generalresearch/managers/thl/user_manager/__init__.py b/generalresearch/managers/thl/user_manager/__init__.py index 0392edb..b3fa8f6 100644 --- a/generalresearch/managers/thl/user_manager/__init__.py +++ b/generalresearch/managers/thl/user_manager/__init__.py @@ -63,7 +63,7 @@ def parse_bp_trust_df(fp: str | Path) -> dict[str, Any]: "entrance_limit_value": convert_int, "median_daily_completes_7d": convert_int, } - bptrust = dict() + bptrust = {} with open(fp, newline="") as csvfile: reader = csv.reader(csvfile) diff --git a/generalresearch/managers/thl/user_manager/mysql_user_manager.py b/generalresearch/managers/thl/user_manager/mysql_user_manager.py index d2d0ffc..e0a7548 100644 --- a/generalresearch/managers/thl/user_manager/mysql_user_manager.py +++ b/generalresearch/managers/thl/user_manager/mysql_user_manager.py @@ -139,11 +139,10 @@ class MysqlUserManager: """) try: - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute(query=query, params=params) - user_id = c.fetchone()["id"] - except psycopg.IntegrityError as e: + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute(query=query, params=params) + user_id = c.fetchone()["id"] + except psycopg.IntegrityError: # Two machines/processes are trying to create this same (product_id, product_user_id) # at the same time. There's a unique index, so mysql will not let two be created. # The 2nd should get an IntegrityError, meaning this already exists, and we can just query it. @@ -160,7 +159,7 @@ class MysqlUserManager: else: # We specifically queried the NON read-replica, and we got an IntegrityError, so # something else must be wrong... - raise e + raise else: user = User( user_id=user_id, diff --git a/generalresearch/managers/thl/userhealth.py b/generalresearch/managers/thl/userhealth.py index babed04..26f08b4 100644 --- a/generalresearch/managers/thl/userhealth.py +++ b/generalresearch/managers/thl/userhealth.py @@ -221,7 +221,7 @@ class IPRecordManager(PostgresManagerWithRedis): "forwarded_ip5", "forwarded_ip6", ] - for col, ip in zip_longest( + for col, fwd_ip in zip_longest( fips_cols, [ forwarded_ip1, @@ -233,7 +233,7 @@ class IPRecordManager(PostgresManagerWithRedis): ], fillvalue=None, ): - data[col] = ipaddress.ip_address(ip).exploded if ip else ip + data[col] = ipaddress.ip_address(fwd_ip).exploded if fwd_ip else fwd_ip self.pg_config.execute_write( query=""" @@ -490,9 +490,7 @@ class AuditLogManager(PostgresManager): created_after: datetime | None = None, ) -> tuple[str, dict[str, Any]]: assert user_ids, "must pass at least 1 user_id" - assert all( - [isinstance(uid, int) for uid in user_ids] - ), "must pass user_id as int" + assert all(isinstance(uid, int) for uid in user_ids), "must pass user_id as int" if created_after is None: created_after = datetime.now(tz=UTC) - timedelta(days=7) diff --git a/generalresearch/managers/thl/wall.py b/generalresearch/managers/thl/wall.py index 03ca1c6..ac9fb62 100644 --- a/generalresearch/managers/thl/wall.py +++ b/generalresearch/managers/thl/wall.py @@ -484,7 +484,7 @@ class WallManager(PostgresManager): ORDER BY rs.source, rs.survey_id; """ - params = dict() + params = {} filters = [] # Instead of doing a big IN with a big set of tuples, since we know diff --git a/generalresearch/managers/thl/wallet/tango.py b/generalresearch/managers/thl/wallet/tango.py index 4abfc70..be8fd97 100644 --- a/generalresearch/managers/thl/wallet/tango.py +++ b/generalresearch/managers/thl/wallet/tango.py @@ -44,7 +44,7 @@ def complete_tango_order( tango_client=tango_client, ) - except Exception: + except AssertionError: # todo: its possible the order went through, but something else was wrong # we should try to retrieve the order by its ref_id and confirm it really # failed... @@ -70,8 +70,8 @@ def create_tango_order( """ Create a tango gift card order. Throws exception if anything is not right. - # https://integration-www.tangocard.com/raas_api_console/v2/ - # https://www.apimatic.io/apidocs/tangocard/v/2_3_4#/python + - https://integration-www.tangocard.com/raas_api_console/v2/ + - https://www.apimatic.io/apidocs/tangocard/v/2_3_4#/python :param utid: Card identifier :param amount: requested card value in USD diff --git a/generalresearch/mariadb.py b/generalresearch/mariadb.py index 8bcd8ee..5d43f97 100644 --- a/generalresearch/mariadb.py +++ b/generalresearch/mariadb.py @@ -32,11 +32,3 @@ def example(): for m in zip(c.metadata["field"], c.metadata["ext_type_or_format"]): # here we can just check if the field's ext_field_flag == 'UUID' (2) print(m[0], ext_field_flags_rev[m[1]]) - - -def get_column_types(): - # How does django do this? - res = """ - SELECT column_name, data_type - FROM information_schema.columns - WHERE table_name = 'morning_userpid' AND table_schema = DATABASE()""" diff --git a/generalresearch/models/admin/__init__.py b/generalresearch/models/admin/__init__.py index ad6302b..344c34a 100644 --- a/generalresearch/models/admin/__init__.py +++ b/generalresearch/models/admin/__init__.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime import pandas as pd from dateutil import relativedelta diff --git a/generalresearch/models/admin/request.py b/generalresearch/models/admin/request.py index 5fdc784..6112786 100644 --- a/generalresearch/models/admin/request.py +++ b/generalresearch/models/admin/request.py @@ -119,7 +119,7 @@ class ReportRequest(BaseModel): @property def end_naive(self) -> datetime: - return datetime.now(tz=None) + return datetime.now(tz=None) # noqa @property def ts_start(self) -> pd.Timestamp: diff --git a/generalresearch/models/cint/question.py b/generalresearch/models/cint/question.py index f8a287a..4c0e52f 100644 --- a/generalresearch/models/cint/question.py +++ b/generalresearch/models/cint/question.py @@ -44,7 +44,7 @@ class CintQuestionType(StrEnum): # This seems to be invalid as there are no options??? "Grid": None, } - return API_TYPE_MAP[a] if a in API_TYPE_MAP else None + return API_TYPE_MAP.get(a) class CintUserQuestionAnswer(MarketplaceUserQuestionAnswer): diff --git a/generalresearch/models/cint/survey.py b/generalresearch/models/cint/survey.py index cfc91ef..fde4559 100644 --- a/generalresearch/models/cint/survey.py +++ b/generalresearch/models/cint/survey.py @@ -12,6 +12,7 @@ from pydantic import ( ConfigDict, Field, NonNegativeInt, + ValidationError, computed_field, model_validator, ) @@ -67,7 +68,7 @@ class CintQuota(BaseModel): condition_hashes: list[str] | None = Field(min_length=1, default=None) def __hash__(self): - return hash(tuple((tuple(self.condition_hashes), self.quota_id))) + return hash((tuple(self.condition_hashes), self.quota_id)) @model_validator(mode="after") def validate_condition_len(self) -> Self: @@ -317,7 +318,7 @@ class CintSurvey(MarketplaceTask): def from_api(cls, d: dict[str, Any]) -> Self | None: try: return cls._from_api(d) - except Exception as e: + except ValidationError as e: logger.warning(f"Unable to parse survey: {d}. {e}") return None @@ -370,8 +371,8 @@ class CintSurvey(MarketplaceTask): d["mobile_conversion"] = None d["revenue_per_click"] = None - d["conditions"] = dict() - d.setdefault("survey_qualifications", list()) + d["conditions"] = {} + d.setdefault("survey_qualifications", []) qualifications = [CintCondition.from_api(q) for q in d["survey_qualifications"]] for q in qualifications: d["conditions"][q.criterion_hash] = q @@ -416,7 +417,7 @@ class CintSurvey(MarketplaceTask): return d @classmethod - def from_mysql(cls, d: Dict[str, Any]) -> Self: + def from_mysql(cls, d: dict[str, Any]) -> Self: d["created_at"] = d["created_at"].replace(tzinfo=UTC) d["last_updated"] = d["last_updated"].replace(tzinfo=UTC) d["qualifications"] = json.loads(d["qualifications"]) @@ -465,7 +466,7 @@ class CintSurvey(MarketplaceTask): ) -> tuple[bool | None, set[str]]: # Many surveys have 0 quotas. Quotas are exclusionary. # They can NOT match a quota where currently_open=0 - total_quota = [q for q in self.quotas if q.quota_type == "total"][0] + total_quota = next(q for q in self.quotas if q.quota_type == "total") if not total_quota.is_open: return False, set() quotas = [q for q in self.quotas if q.quota_type != "total"] @@ -474,7 +475,7 @@ class CintSurvey(MarketplaceTask): quota_eval = { quota: quota.matches_soft(criteria_evaluation) for quota in quotas } - evals = set(g[0] for g in quota_eval.values()) + evals = {g[0] for g in quota_eval.values()} if any(m[0] is True and not q.is_open for q, m in quota_eval.items()): # matched a full quota return False, set() diff --git a/generalresearch/models/custom_types.py b/generalresearch/models/custom_types.py index c200b34..5e4db3e 100644 --- a/generalresearch/models/custom_types.py +++ b/generalresearch/models/custom_types.py @@ -98,7 +98,7 @@ LanguageISOLike = Annotated[ def check_valid_uuid(v: str) -> str: try: assert UUID(v).hex == v - except Exception: + except (ValueError, AssertionError): raise ValueError("Invalid UUID") return v @@ -106,7 +106,7 @@ def check_valid_uuid(v: str) -> str: def is_valid_uuid(v: str) -> bool: try: assert UUID(v).hex == v - except Exception: + except (ValueError, AssertionError): return False return True @@ -165,7 +165,7 @@ CoercedStr = Annotated[str, BeforeValidator(coerce_int_to_str)] # Serializers that can transform a collection of str into a comma separated # str bidirectionally -to_comma_sep_str = PlainSerializer(lambda x: ",".join(sorted(list(x))), return_type=str) +to_comma_sep_str = PlainSerializer(lambda x: ",".join(sorted(x)), return_type=str) enum_to_comma_sep_str = PlainSerializer( lambda x: ",".join(sorted([str(y.value) for y in x])), return_type=str ) diff --git a/generalresearch/models/dynata/survey.py b/generalresearch/models/dynata/survey.py index 5a9f763..942ab4f 100644 --- a/generalresearch/models/dynata/survey.py +++ b/generalresearch/models/dynata/survey.py @@ -168,7 +168,7 @@ class DynataQuota(BaseModel): status: DynataStatus = Field() def __hash__(self): - return hash(tuple((tuple(self.condition_hashes), self.count, self.status))) + return hash((tuple(self.condition_hashes), self.count, self.status)) @property def is_open(self) -> bool: @@ -244,7 +244,7 @@ class DynataQuotaGroup(RootModel): ) -> tuple[bool | None, set[str]]: # Qualify for ANY quota object within a quota group obj_evals = {obj: obj.passes_soft(criteria_evaluation) for obj in self.root} - evals = set(v[0] for v in obj_evals.values()) + evals = {v[0] for v in obj_evals.values()} # If we match 1 obj, then the others don't matter if any(evals): return True, set() @@ -319,7 +319,7 @@ class DynataFilterGroup(RootModel): ) -> tuple[bool | None, set[str]]: # Passes back "passes" (T/F/none) and a list of unknown criterion hashes obj_evals = {obj: obj.passes_soft(criteria_evaluation) for obj in self.root} - evals = set(v[0] for v in obj_evals.values()) + evals = {v[0] for v in obj_evals.values()} # If we match 1 obj, then the others don't matter if any(evals): return True, set() @@ -548,7 +548,7 @@ class DynataSurvey(MarketplaceTask): return d @classmethod - def from_db(cls, d: Dict[str, Any]) -> Self: + def from_db(cls, d: dict[str, Any]) -> Self: d["created"] = d["created"].replace(tzinfo=UTC) d["last_updated"] = d["last_updated"].replace(tzinfo=UTC) d["filters"] = json.loads(d["filters"]) @@ -578,7 +578,7 @@ class DynataSurvey(MarketplaceTask): group_eval = { group: group.passes_soft(criteria_evaluation) for group in self.filters } - evals = set(g[0] for g in group_eval.values()) + evals = {g[0] for g in group_eval.values()} if False in evals: return False, set() elif None in evals: @@ -614,7 +614,7 @@ class DynataSurvey(MarketplaceTask): group_eval = { quota: quota.passes_soft(criteria_evaluation) for quota in self.quotas } - evals = set(g[0] for g in group_eval.values()) + evals = {g[0] for g in group_eval.values()} if False in evals: return False, set() elif None in evals: diff --git a/generalresearch/models/dynata/task_collection.py b/generalresearch/models/dynata/task_collection.py index 71cf3db..2b82bfd 100644 --- a/generalresearch/models/dynata/task_collection.py +++ b/generalresearch/models/dynata/task_collection.py @@ -54,7 +54,7 @@ DynataTaskCollectionSchema = DataFrameSchema( class DynataTaskCollection(TaskCollection): - items: List[DynataSurvey] + items: list[DynataSurvey] _schema = DynataTaskCollectionSchema def to_row(self, s: DynataSurvey) -> dict[str, Any]: diff --git a/generalresearch/models/gr/authentication.py b/generalresearch/models/gr/authentication.py index f9644fe..25f65fa 100644 --- a/generalresearch/models/gr/authentication.py +++ b/generalresearch/models/gr/authentication.py @@ -4,7 +4,7 @@ import binascii import json import os from datetime import UTC, datetime -from typing import TYPE_CHECKING, Any, Self +from typing import TYPE_CHECKING, Any from pydantic import ( AnyHttpUrl, @@ -283,16 +283,15 @@ class GRUser(BaseModel): ex=ex_secs, ) - # --- ORM --- @classmethod - def from_postgresql(cls, d: dict) -> Self: + def from_postgresql(cls, d: dict[str, Any]) -> GRUser: d["date_joined"] = d["date_joined"].replace(tzinfo=UTC) return GRUser.model_validate(d) @classmethod - def from_redis(cls, d: str | dict[str, Any]) -> Self: + def from_redis(cls, d: str | dict[str, Any]) -> GRUser: if isinstance(d, str): d = json.loads(d) assert isinstance(d, dict) @@ -357,13 +356,13 @@ class GRToken(BaseModel): # --- Properties --- @property - def auth_header(self, key_name: str = "Authorization") -> dict[str, str]: - return {key_name: self.key} + def auth_header(self) -> dict[str, str]: + return {"Authorization": self.key} # --- ORM --- @classmethod - def from_redis(cls, d: str | dict[str, Any]) -> Self: + def from_redis(cls, d: str | dict[str, Any]) -> GRToken: if isinstance(d, str): d = json.loads(d) assert isinstance(d, dict) diff --git a/generalresearch/models/gr/business.py b/generalresearch/models/gr/business.py index 74b5c29..064c200 100644 --- a/generalresearch/models/gr/business.py +++ b/generalresearch/models/gr/business.py @@ -6,14 +6,15 @@ import os from datetime import UTC, datetime from enum import Enum, StrEnum from pathlib import Path -from typing import TYPE_CHECKING, Self +from typing import TYPE_CHECKING from uuid import uuid4 import pandas as pd +import pyarrow as pa from dask.distributed import Client from psycopg.cursor import Cursor from psycopg.rows import dict_row -from pydantic import BaseModel, ConfigDict, Field, PositiveInt +from pydantic import BaseModel, ConfigDict, Field, PositiveInt, ValidationError from pydantic.json_schema import SkipJsonSchema from pydantic_extra_types.phone_numbers import PhoneNumber @@ -210,9 +211,11 @@ class Business(BaseModel): payouts: list[BusinessPayoutEvent] | None = Field( default=None, name="Business Payouts", - description="These are the ACH or Wire payments that were sent to the" - "Business as a single amount, summed for all the Business" - "child Products", + description=( + "These are the ACH or Wire payments that were sent to the" + "Business as a single amount, summed for all the Business" + "child Products" + ), ) pop_financial: list[POPFinancial] | None = Field(default=None) @@ -237,18 +240,19 @@ class Business(BaseModel): # --- Prefetch --- def prefetch_addresses(self, pg_config: PostgresConfig) -> None: - with pg_config.make_connection() as conn: - with conn.cursor(row_factory=dict_row) as c: - c.execute( - query=""" + with pg_config.make_connection() as conn, conn.cursor( + row_factory=dict_row + ) as c: + c.execute( + query=""" SELECT * FROM common_businessaddress AS ba WHERE ba.business_id = %s LIMIT 1 """, - params=[self.id], - ) - res = c.fetchall() + params=[self.id], + ) + res = c.fetchall() if len(res) == 0: self.addresses = [] @@ -258,22 +262,23 @@ class Business(BaseModel): def prefetch_teams(self, pg_config: PostgresConfig) -> None: from generalresearch.models.gr.team import Team - with pg_config.make_connection() as conn: - with conn.cursor(row_factory=dict_row) as c: - c: Cursor + with pg_config.make_connection() as conn, conn.cursor( + row_factory=dict_row + ) as c: + c: Cursor - c.execute( - query=""" + c.execute( + query=""" SELECT t.* FROM common_team AS t INNER JOIN common_team_businesses AS tb ON tb.team_id = t.id WHERE tb.business_id = %s """, - params=(self.id,), - ) + params=(self.id,), + ) - res = c.fetchall() + res = c.fetchall() if len(res) == 0: self.teams = [] @@ -542,11 +547,10 @@ class Business(BaseModel): ) try: - test = pd.read_parquet(path, engine="pyarrow") - except Exception as e: + _ = pd.read_parquet(path, engine="pyarrow") + except (pa.ArrowException, OSError, ValueError) as e: raise OSError(f"Parquet verification failed: {e}") - def prebuild_enriched_wall_parquet( self, thl_pg_config: PostgresConfig, @@ -586,11 +590,10 @@ class Business(BaseModel): ) try: - test = pd.read_parquet(path, engine="pyarrow") - except Exception as e: + _ = pd.read_parquet(path, engine="pyarrow") + except (pa.ArrowException, OSError, ValueError) as e: raise OSError(f"Parquet verification failed: {e}") - @classmethod def required_fields(cls) -> list[str]: return [ @@ -651,7 +654,7 @@ class Business(BaseModel): client=client, pop_ledger=pop_ledger, ) - self.prebuild_payouts(thl_pg_config=thl_web_rr, thl_lm=thl_lm, bpem=bpem) + self.prebuild_payouts(bpem=bpem) self.prebuild_pop_financial( thl_pg_config=thl_web_rr, thl_lm=thl_lm, @@ -713,7 +716,7 @@ class Business(BaseModel): uuid: UUIDStr, fields: list[str], gr_redis_config: RedisConfig, - ) -> Self | None: + ) -> Business | None: keys: list[str] = Business.required_fields() + fields if "pop_financial" in keys: @@ -724,7 +727,7 @@ class Business(BaseModel): rc = gr_redis_config.create_redis_client() try: - res: list = rc.hmget(name=f"business:{uuid}", keys=keys) + res: list[str | bytes | None] = rc.hmget(name=f"business:{uuid}", keys=keys) d = { val: json.loads(res[idx]) if res[idx] is not None else None for idx, val in enumerate(keys) @@ -742,6 +745,5 @@ class Business(BaseModel): result["pop_financial"] = pop_financial return Business.model_validate(result) - except Exception as e: - logging.exception(e) + except ValidationError: return None diff --git a/generalresearch/models/gr/team.py b/generalresearch/models/gr/team.py index 78a9ba9..4752bea 100644 --- a/generalresearch/models/gr/team.py +++ b/generalresearch/models/gr/team.py @@ -5,16 +5,18 @@ import os from datetime import UTC, datetime from enum import Enum from pathlib import Path -from typing import TYPE_CHECKING, Self +from typing import TYPE_CHECKING from uuid import uuid4 import pandas as pd +import pyarrow as pa from dask.distributed import Client from pydantic import ( BaseModel, ConfigDict, Field, PositiveInt, + ValidationError, field_validator, ) from pydantic.json_schema import SkipJsonSchema @@ -191,10 +193,9 @@ class Team(BaseModel): try: _ = pd.read_parquet(path, engine="pyarrow") - except Exception as e: + except (pa.ArrowException, OSError, ValueError) as e: raise OSError(f"Parquet verification failed: {e}") - def prebuild_enriched_wall_parquet( self, thl_pg_config: PostgresConfig, @@ -235,10 +236,9 @@ class Team(BaseModel): try: _ = pd.read_parquet(path, engine="pyarrow") - except Exception as e: + except (pa.ArrowException, OSError, ValueError) as e: raise OSError(f"Parquet verification failed: {e}") - @classmethod def required_fields(cls) -> list[str]: return [ @@ -281,8 +281,6 @@ class Team(BaseModel): enriched_session: EnrichedSessionMerge | None = None, enriched_wall: EnrichedWallMerge | None = None, ) -> None: - ex_secs = 60 * 60 * 24 * 3 # 3 days - self.prefetch_products(thl_pg_config=thl_web_rr) self.prefetch_gr_users(pg_config=pg_config, redis_config=redis_config) self.prefetch_businesses(pg_config=pg_config, redis_config=redis_config) @@ -323,7 +321,6 @@ class Team(BaseModel): enriched_wall=enriched_wall, ) - # --- ORM --- @classmethod @@ -332,14 +329,14 @@ class Team(BaseModel): uuid: UUIDStr, fields: list[str], gr_redis_config: RedisConfig, - ) -> Self | None: + ) -> Team | None: keys: list = Team.required_fields() + fields rc = gr_redis_config.create_redis_client() try: - res: list = rc.hmget(name=f"team:{uuid}", keys=keys) + res: list[str | bytes | None] = rc.hmget(name=f"team:{uuid}", keys=keys) d = {val: json.loads(res[idx]) for idx, val in enumerate(keys)} return Team.model_validate(d) - except Exception: + except ValidationError: return None diff --git a/generalresearch/models/innovate/question.py b/generalresearch/models/innovate/question.py index 89310a2..f5a4846 100644 --- a/generalresearch/models/innovate/question.py +++ b/generalresearch/models/innovate/question.py @@ -6,7 +6,7 @@ import logging from enum import StrEnum from typing import TYPE_CHECKING, Any, Literal -from pydantic import BaseModel, Field, field_validator, model_validator +from pydantic import BaseModel, Field, ValidationError, field_validator, model_validator from generalresearch.models import Source from generalresearch.models.innovate import InnovateQuestionID @@ -71,7 +71,7 @@ class InnovateQuestionType(StrEnum): @classmethod def from_api(cls, a: int): API_TYPE_MAP = cls.get_api_map() - return API_TYPE_MAP[a] if a in API_TYPE_MAP else None + return API_TYPE_MAP.get(a) class InnovateQuestion(MarketplaceQuestion): @@ -141,7 +141,7 @@ class InnovateQuestion(MarketplaceQuestion): @classmethod def from_api( - cls, d: dict, country_iso: str, language_iso: str + cls, d: dict[str, Any], country_iso: str, language_iso: str ) -> InnovateQuestion | None: """ :param d: Raw response from API @@ -151,13 +151,13 @@ class InnovateQuestion(MarketplaceQuestion): """ try: return cls._from_api(d, country_iso, language_iso) - except Exception as e: + except ValidationError as e: logger.warning(f"Unable to parse question: {d}. {e}") return None @classmethod def _from_api( - cls, d: dict, country_iso: str, language_iso: str + cls, d: dict[str, Any], country_iso: str, language_iso: str ) -> InnovateQuestion: # Question AGE returns options even though its marked as a text entry (but only in some locales) d["QuestionKey"] = d["QuestionKey"].lower() diff --git a/generalresearch/models/innovate/survey.py b/generalresearch/models/innovate/survey.py index 3c37fe3..d07f960 100644 --- a/generalresearch/models/innovate/survey.py +++ b/generalresearch/models/innovate/survey.py @@ -17,6 +17,7 @@ from pydantic import ( BaseModel, ConfigDict, Field, + ValidationError, computed_field, model_validator, ) @@ -69,7 +70,7 @@ class InnovateCondition(MarketplaceCondition): d["logical_operator"] = LogicalOperator.OR d["value_type"] = ConditionValueType.LIST d["negate"] = False - d["values"] = list(set(x.strip().lower() for x in d["values"])) + d["values"] = list({x.strip().lower() for x in d["values"]}) return cls.model_validate(d) @@ -88,7 +89,7 @@ class InnovateQuota(BaseModel): condition_hashes: list[str] = Field(min_length=0, default_factory=list) def __hash__(self): - return hash(tuple((tuple(self.condition_hashes), self.remaining_count))) + return hash((tuple(self.condition_hashes), self.remaining_count)) @property def is_open(self) -> bool: @@ -99,7 +100,7 @@ class InnovateQuota(BaseModel): ) @classmethod - def from_api(cls, d: dict): + def from_api(cls, d: dict[str, Any]): return cls.model_validate(d) def passes(self, criteria_evaluation: dict[str, bool | None]) -> bool: @@ -263,13 +264,13 @@ class InnovateSurvey(MarketplaceTask): def from_api(cls, d: dict[str, Any]) -> InnovateSurvey | None: try: return cls._from_api(d) - except Exception as e: + except ValidationError as e: logger.warning(f"Unable to parse survey: {d}. {e}") return None @classmethod def _from_api(cls, d: dict[str, Any]) -> InnovateSurvey: - d["conditions"] = dict() + d["conditions"] = {} # If we haven't hit the "detail" endpoint, we won't get this d.setdefault("qualifications", []) @@ -317,11 +318,14 @@ class InnovateSurvey(MarketplaceTask): # Fancy repr that abbreviates exclude_pids and excluded_surveys repr_args = list(self.__repr_args__()) for n, (k, v) in enumerate(repr_args): - if k in {"exclude_pids", "include_pids", "excluded_surveys"}: - if v and len(v) > 6: - v = sorted(v) - v = v[:3] + ["…"] + v[-3:] - repr_args[n] = (k, v) + if ( + k in {"exclude_pids", "include_pids", "excluded_surveys"} + and v + and len(v) > 6 + ): + v = sorted(v) + v = v[:3] + ["…"] + v[-3:] + repr_args[n] = (k, v) join_str = ", " repr_str = join_str.join( repr(v) if a is None else f"{a}={v!r}" for a, v in repr_args @@ -380,14 +384,21 @@ class InnovateSurvey(MarketplaceTask): """ assert isinstance(att_survey_ids, set), "must pass a set" assert isinstance(att_job_ids, set), "must pass a set" + if self.survey_id in att_survey_ids: return False - if self.duplicate_check_level == InnovateDuplicateCheckLevel.JOB: - if self.job_id in att_job_ids: - return False + + if ( + self.duplicate_check_level == InnovateDuplicateCheckLevel.JOB + and self.job_id in att_job_ids + ): + return False + if self.duplicate_check_level == InnovateDuplicateCheckLevel.EXCLUDED_SURVEYS: + assert self.excluded_surveys is not None if self.excluded_surveys.intersection(att_survey_ids): return False + return True def passes_qualifications( @@ -431,7 +442,7 @@ class InnovateSurvey(MarketplaceTask): quota_eval = { quota: quota.matches_soft(criteria_evaluation) for quota in self.quotas } - evals = set(g[0] for g in quota_eval.values()) + evals = {g[0] for g in quota_eval.values()} if any(m[0] is True and not q.is_open for q, m in quota_eval.items()): # matched a full quota return False, set() diff --git a/generalresearch/models/legacy/bucket.py b/generalresearch/models/legacy/bucket.py index 812241d..3705f4f 100644 --- a/generalresearch/models/legacy/bucket.py +++ b/generalresearch/models/legacy/bucket.py @@ -120,8 +120,10 @@ class BucketBase(BaseModel): ) uri: HttpsUrl = Field( examples=[ - "https://task.generalresearch.com/api/v1/52d3f63b2709/797df4136c604a6c8599818296aae6d1/?i" - "=5ba2fe5010cc4d078fc3cc0b0cc264c3&b=test&66482fb=e7baf5e" + ( + "https://task.generalresearch.com/api/v1/52d3f63b2709/797df4136c604a6c8599818296aae6d1/?i" + "=5ba2fe5010cc4d078fc3cc0b0cc264c3&b=test&66482fb=e7baf5e" + ) ], description="The URL to send a respondent into. Must not edit this URL in any way", ) @@ -465,12 +467,12 @@ class PayoutSummaryDecimal(StatisticalSummary): class PayoutSummary(StatisticalSummary): """Payouts are in Integer USD Cents""" - min: int = Field(gt=0, le=10000) - max: int = Field(gt=0, le=10000) - q1: int = Field(gt=0, le=10000) - q2: int = Field(gt=0, le=10000) - q3: int = Field(gt=0, le=10000) - mean: int | None = Field(gt=0, le=10000, default=None) + min: int = Field(gt=0, le=10_000) + max: int = Field(gt=0, le=10_000) + q1: int = Field(gt=0, le=10_000) + q2: int = Field(gt=0, le=10_000) + q3: int = Field(gt=0, le=10_000) + mean: int | None = Field(gt=0, le=10_000, default=None) model_config = { "json_schema_extra": { @@ -724,8 +726,10 @@ class OneShotOfferwallBucket(BaseModel): ) uri: HttpsUrl = Field( examples=[ - "https://task.generalresearch.com/api/v1/52d3f63b2709/797df4136c604a6c8599818296aae6d1/?i" - "=5ba2fe5010cc4d078fc3cc0b0cc264c3&b=test&66482fb=e7baf5e" + ( + "https://task.generalresearch.com/api/v1/52d3f63b2709/797df4136c604a6c8599818296aae6d1/?i" + "=5ba2fe5010cc4d078fc3cc0b0cc264c3&b=test&66482fb=e7baf5e" + ) ], description="The URL to send a respondent into. Must not edit this URL in any way", ) @@ -759,8 +763,10 @@ class WXETOfferwallBucket(BaseModel): ) uri: HttpsUrl = Field( examples=[ - "https://task.generalresearch.com/api/v1/52d3f63b2709/797df4136c604a6c8599818296aae6d1/?i" - "=5ba2fe5010cc4d078fc3cc0b0cc264c3&b=test&66482fb=e7baf5e" + ( + "https://task.generalresearch.com/api/v1/52d3f63b2709/797df4136c604a6c8599818296aae6d1/?i" + "=5ba2fe5010cc4d078fc3cc0b0cc264c3&b=test&66482fb=e7baf5e" + ) ], description="The URL to send a respondent into. Must not edit this URL in any way", ) diff --git a/generalresearch/models/legacy/questions.py b/generalresearch/models/legacy/questions.py index 4651ab0..9f37837 100644 --- a/generalresearch/models/legacy/questions.py +++ b/generalresearch/models/legacy/questions.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Annotated, Any, Self +from typing import TYPE_CHECKING, Annotated, Any from pydantic import ( BaseModel, @@ -87,25 +87,17 @@ class UserQuestionAnswerIn(BaseModel): fingerprint_tz = "a91cb1dea814480dba12d9b7b48696dd" fingerprint_fingerprint = "1d1e2e8380ac474b87fb4e4c569b48df" - if self.question_id in { - user_agent_qid, - fingerprint_langs, - fingerprint_tz, - fingerprint_fingerprint, - }: - if len(self.answer) != 1: - raise ValueError("Too many answer values provided") - - return self - - @model_validator(mode="after") - def user_agent_check(self) -> Self: - # TODO: where / how do I want to pass in this Werz user_agent stuff? - user_agent_qid = "2fbedb2b9f7647b09ff5e52fa119cc5e" - - if self.question_id == user_agent_qid: - val = self.answer[0] - # assert val == request.user_agent.to_header(): + if ( + self.question_id + in { + user_agent_qid, + fingerprint_langs, + fingerprint_tz, + fingerprint_fingerprint, + } + and len(self.answer) != 1 + ): + raise ValueError("Too many answer values provided") return self diff --git a/generalresearch/models/morning/survey.py b/generalresearch/models/morning/survey.py index 5011698..91d1bce 100644 --- a/generalresearch/models/morning/survey.py +++ b/generalresearch/models/morning/survey.py @@ -200,7 +200,7 @@ class MorningQuota(MorningStatistics, MarketplaceTask): data["country_isos"] = [data["country_iso"]] if isinstance(data["language_isos"], str): data["language_isos"] = set(data["language_isos"].split(",")) - data["language_iso"] = sorted(data["language_isos"])[0] + data["language_iso"] = min(data["language_isos"]) return data @property @@ -276,11 +276,11 @@ class MorningQuota(MorningStatistics, MarketplaceTask): self, criteria_evaluation: dict[str, bool | None] ) -> tuple[bool | None, list[str]]: # Passes back "matches" (T/F/none) and a list of unknown criterion hashes - unknowns = list() + unknowns = [] for c in self.condition_hashes: eval_value = criteria_evaluation.get(c) if eval_value is False: - return False, list() + return False, [] if eval_value is None: unknowns.append(c) if unknowns: @@ -359,7 +359,7 @@ class MorningBid(MorningTaskStatistics): @property def language_iso_any(self): - return sorted(self.language_isos)[0] + return min(self.language_isos) @property def locale(self): @@ -417,7 +417,7 @@ class MorningBid(MorningTaskStatistics): if "conditions" in data: return data - data["conditions"] = dict() + data["conditions"] = {} for quota in data["quotas"]: if "qualifications" in quota: quota_conditions = [ diff --git a/generalresearch/models/morning/task_collection.py b/generalresearch/models/morning/task_collection.py index eb4cbd1..9303a2f 100644 --- a/generalresearch/models/morning/task_collection.py +++ b/generalresearch/models/morning/task_collection.py @@ -108,7 +108,7 @@ class MorningTaskCollection(TaskCollection): ] quota_fields = list(quota_columns.keys()) rows = [] - bid_dict = dict() + bid_dict = {} for k in bid_fields: bid_dict[k] = getattr(bid, k) bid_dict["bid.id"] = bid.id diff --git a/generalresearch/models/network/label.py b/generalresearch/models/network/label.py index e4ddd18..60a6e58 100644 --- a/generalresearch/models/network/label.py +++ b/generalresearch/models/network/label.py @@ -2,6 +2,7 @@ from __future__ import annotations import ipaddress from enum import StrEnum +from ipaddress import IPv4Network, IPv6Network from pydantic import ( BaseModel, @@ -84,12 +85,13 @@ class IPLabel(BaseModel): @field_validator("ip", mode="before") @classmethod - def normalize_and_validate_network(cls, v): - net = ipaddress.ip_network(v, strict=False) + def normalize_and_validate_network( + cls, v: IPvAnyNetwork + ) -> IPv4Network | IPv6Network | None: + net = ipaddress.ip_network(address=v, strict=False) - if isinstance(net, ipaddress.IPv6Network): - if net.prefixlen > 64: - raise ValueError("IPv6 network must be /64 or larger") + if isinstance(net, ipaddress.IPv6Network) and net.prefixlen > 64: + raise ValueError("IPv6 network must be /64 or larger") return net diff --git a/generalresearch/models/network/nmap/result.py b/generalresearch/models/network/nmap/result.py index 55c2109..4552e15 100644 --- a/generalresearch/models/network/nmap/result.py +++ b/generalresearch/models/network/nmap/result.py @@ -411,7 +411,7 @@ class NmapResult(BaseModel): def model_dump_postgres(self): # Writes for the network_portscan table - d = dict() + d = {} data = self.model_dump( mode="json", include={ diff --git a/generalresearch/models/network/rdns/command.py b/generalresearch/models/network/rdns/command.py index e88a84d..bccead0 100644 --- a/generalresearch/models/network/rdns/command.py +++ b/generalresearch/models/network/rdns/command.py @@ -20,7 +20,7 @@ def run_rdns(config: RDNSRunCommand) -> RDNSResult: def build_rdns_command(ip: str) -> str: # e.g. dig +noall +answer -x 1.2.3.4 - return " ".join(["dig", "+noall", "+answer", "-x", ip]) + return f"dig +noall +answer -x {ip}" def get_dig_version() -> str: diff --git a/generalresearch/models/precision/question.py b/generalresearch/models/precision/question.py index cc90aa9..f532998 100644 --- a/generalresearch/models/precision/question.py +++ b/generalresearch/models/precision/question.py @@ -54,15 +54,15 @@ class PrecisionQuestionType(StrEnum): TEXT_ENTRY = "t" @classmethod - def from_api(cls, a: int): - API_TYPE_MAP = { + def from_api(cls, a: int) -> PrecisionQuestionType | None: + api_type_map: dict[str, PrecisionQuestionType] = { "Drop Down": PrecisionQuestionType.SINGLE_SELECT, "Multi Select": PrecisionQuestionType.MULTI_SELECT, "Single Select": PrecisionQuestionType.SINGLE_SELECT, "Single Select Matrix": PrecisionQuestionType.SINGLE_SELECT, "Vertical Question": PrecisionQuestionType.SINGLE_SELECT, } - return API_TYPE_MAP[a] if a in API_TYPE_MAP else None + return api_type_map.get(a, None) class PrecisionUserQuestionAnswer(MarketplaceUserQuestionAnswer): diff --git a/generalresearch/models/precision/survey.py b/generalresearch/models/precision/survey.py index b27b8c4..bf9e83e 100644 --- a/generalresearch/models/precision/survey.py +++ b/generalresearch/models/precision/survey.py @@ -312,7 +312,7 @@ class PrecisionSurvey(MarketplaceTask): quota_eval = { quota: quota.matches_soft(criteria_evaluation) for quota in self.quotas } - evals = set(g[0] for g in quota_eval.values()) + evals = {g[0] for g in quota_eval.values()} if any(m[0] is True and not q.is_open for q, m in quota_eval.items()): # matched a full quota return False, set() diff --git a/generalresearch/models/prodege/question.py b/generalresearch/models/prodege/question.py index 574c4fd..58aed67 100644 --- a/generalresearch/models/prodege/question.py +++ b/generalresearch/models/prodege/question.py @@ -8,7 +8,14 @@ from enum import StrEnum from functools import cached_property from typing import TYPE_CHECKING, Any, Literal -from pydantic import BaseModel, ConfigDict, Field, PositiveInt, model_validator +from pydantic import ( + BaseModel, + ConfigDict, + Field, + PositiveInt, + ValidationError, + model_validator, +) from generalresearch.locales import Localelator from generalresearch.models import MAX_INT32, Source @@ -143,7 +150,7 @@ class ProdegeQuestion(MarketplaceQuestion): """ try: return cls._from_api(d, country_iso) - except Exception as e: + except ValidationError as e: logger.warning(f"Unable to parse question: {d}. {e}") return None diff --git a/generalresearch/models/prodege/survey.py b/generalresearch/models/prodege/survey.py index 7ab6df6..5d0369a 100644 --- a/generalresearch/models/prodege/survey.py +++ b/generalresearch/models/prodege/survey.py @@ -539,7 +539,7 @@ class ProdegeSurvey(MarketplaceTask): d["country_isos"] = [ locale_helper.get_country_iso(d.pop("country_code").lower()) ] - d["country_iso"] = sorted(d["country_isos"])[0] + d["country_iso"] = min(d["country_isos"]) # No languages are returned anywhere for anything d["language_isos"] = [ locale_helper.get_default_lang_from_country(d["country_isos"][0]) @@ -552,7 +552,7 @@ class ProdegeSurvey(MarketplaceTask): d["past_participation"] = ProdegePastParticipation.from_api( d["past_participation"] ) - d["conditions"] = dict() + d["conditions"] = {} for quota in d["quotas"]: quota["condition_hashes"] = [] for c in quota["targeting_criteria"]: @@ -563,7 +563,7 @@ class ProdegeSurvey(MarketplaceTask): d["quotas"] = [ProdegeQuota.from_api(q) for q in d["quotas"]] countries = {q.country_iso for q in d["quotas"] if q.country_iso} if countries: - d["country_iso"] = sorted(countries)[0] + d["country_iso"] = min(countries) d["country_isos"] = countries d["language_iso"] = locale_helper.get_default_lang_from_country( d["country_iso"] diff --git a/generalresearch/models/prodege/task_collection.py b/generalresearch/models/prodege/task_collection.py index 4544050..9f6a81b 100644 --- a/generalresearch/models/prodege/task_collection.py +++ b/generalresearch/models/prodege/task_collection.py @@ -76,7 +76,7 @@ class ProdegeTaskCollection(TaskCollection): "used_question_ids", "all_hashes", ] - d = dict() + d = {} for k in fields: d[k] = getattr(s, k) d["cpi"] = float(d["cpi"]) diff --git a/generalresearch/models/repdata/question.py b/generalresearch/models/repdata/question.py index 9dda97f..4fa2d22 100644 --- a/generalresearch/models/repdata/question.py +++ b/generalresearch/models/repdata/question.py @@ -12,6 +12,7 @@ from pydantic import ( ConfigDict, Field, PositiveInt, + ValidationError, field_validator, model_validator, ) @@ -142,6 +143,7 @@ class RepDataQuestion(MarketplaceQuestion): @property def internal_id(self) -> str: + assert self.lucid_id return self.lucid_id @field_validator("question_id", mode="before") @@ -167,7 +169,7 @@ class RepDataQuestion(MarketplaceQuestion): """ try: return cls._from_api(d, country_iso, language_iso) - except Exception as e: + except ValidationError as e: logger.warning(f"Unable to parse question: {d}. {e}") return None diff --git a/generalresearch/models/repdata/survey.py b/generalresearch/models/repdata/survey.py index 43a592c..c5b0730 100644 --- a/generalresearch/models/repdata/survey.py +++ b/generalresearch/models/repdata/survey.py @@ -13,6 +13,7 @@ from pydantic import ( BaseModel, ConfigDict, Field, + ValidationError, computed_field, field_validator, model_validator, @@ -459,7 +460,7 @@ class RepDataSurvey(BaseModel): @property def all_conditions(self) -> list[RepDataCondition]: - cs = list() + cs = [] for stream in self.streams: cs.extend(stream.all_conditions) # dedupe by criterion_hash @@ -477,7 +478,7 @@ class RepDataSurvey(BaseModel): """ try: return cls._from_api(survey_response) - except Exception as e: + except ValidationError as e: survey_id = survey_response.get("survey_id") or survey_response.get( "SurveyNumber" ) @@ -485,7 +486,7 @@ class RepDataSurvey(BaseModel): return None @classmethod - def _from_api(cls, survey_response) -> RepDataSurvey: + def _from_api(cls, survey_response: dict[str, Any]) -> RepDataSurvey: d = survey_response.copy() d["country_iso"] = locale_helper.get_country_iso(d["SurveyCountry"].lower()) d["language_iso"] = locale_helper.get_language_iso(d["SurveyLanguage"].lower()) diff --git a/generalresearch/models/repdata/task_collection.py b/generalresearch/models/repdata/task_collection.py index d625349..5b9a4ba 100644 --- a/generalresearch/models/repdata/task_collection.py +++ b/generalresearch/models/repdata/task_collection.py @@ -110,7 +110,7 @@ class RepDataTaskCollection(TaskCollection): "remaining_count", ] rows = [] - d = dict() + d = {} for k in survey_fields: d[k] = getattr(s, k) d["allowed_devices"] = s.allowed_devices_str diff --git a/generalresearch/models/sago/question.py b/generalresearch/models/sago/question.py index 474543d..291214f 100644 --- a/generalresearch/models/sago/question.py +++ b/generalresearch/models/sago/question.py @@ -13,6 +13,7 @@ from pydantic import ( ConfigDict, Field, PositiveInt, + ValidationError, field_validator, model_validator, ) @@ -86,7 +87,7 @@ class SagoQuestionType(StrEnum): 6: SagoQuestionType.TEXT_ENTRY, 7: SagoQuestionType.TEXT_ENTRY, } - return API_TYPE_MAP[a] if a in API_TYPE_MAP else None + return API_TYPE_MAP.get(a, None) class SagoUserQuestionAnswer(BaseModel): @@ -182,7 +183,7 @@ class SagoQuestion(MarketplaceQuestion): """ try: return cls._from_api(d, country_iso, language_iso) - except Exception as e: + except ValidationError as e: logger.warning(f"Unable to parse question: {d}. {e}") return None diff --git a/generalresearch/models/sago/survey.py b/generalresearch/models/sago/survey.py index 83aad8c..8550cd3 100644 --- a/generalresearch/models/sago/survey.py +++ b/generalresearch/models/sago/survey.py @@ -8,7 +8,14 @@ from functools import cached_property from typing import Annotated, Any, Literal, Self from more_itertools import flatten -from pydantic import BaseModel, ConfigDict, Field, computed_field, model_validator +from pydantic import ( + BaseModel, + ConfigDict, + Field, + ValidationError, + computed_field, + model_validator, +) from generalresearch.locales import Localelator from generalresearch.models import LogicalOperator, Source @@ -71,7 +78,7 @@ class SagoQuota(BaseModel): # There is no explicit status. The quota is closed if the count is 0 def __hash__(self) -> int: - return hash(tuple((tuple(self.condition_hashes), self.remaining_count))) + return hash((tuple(self.condition_hashes), self.remaining_count)) @property def is_open(self) -> bool: @@ -261,7 +268,7 @@ class SagoSurvey(MarketplaceTask): def from_api(cls, d: dict[str, Any]) -> SagoSurvey | None: try: return cls._from_api(d) - except Exception as e: + except ValidationError as e: logger.warning(f"Unable to parse survey: {d}. {e}") return None @@ -273,11 +280,10 @@ class SagoSurvey(MarketplaceTask): # Fancy repr that abbreviates ip_exclusions and survey_exclusions repr_args = list(self.__repr_args__()) for n, (k, v) in enumerate(repr_args): - if k in {"ip_exclusions", "survey_exclusions"}: - if v and len(v) > 6: - v = sorted(v) - v = v[:3] + ["…"] + v[-3:] - repr_args[n] = (k, v) + if k in {"ip_exclusions", "survey_exclusions"} and v and len(v) > 6: + v = sorted(v) + v = v[:3] + ["…"] + v[-3:] + repr_args[n] = (k, v) join_str = ", " repr_str = join_str.join( repr(v) if a is None else f"{a}={v!r}" for a, v in repr_args @@ -362,7 +368,7 @@ class SagoSurvey(MarketplaceTask): quota_eval = { quota: quota.matches_soft(criteria_evaluation) for quota in self.quotas } - evals = set(g[0] for g in quota_eval.values()) + evals = {g[0] for g in quota_eval.values()} if any(m[0] is True and not q.is_open for q, m in quota_eval.items()): # matched a full quota return False, set() diff --git a/generalresearch/models/thl/contest/__init__.py b/generalresearch/models/thl/contest/__init__.py index 842d85e..65b28f4 100644 --- a/generalresearch/models/thl/contest/__init__.py +++ b/generalresearch/models/thl/contest/__init__.py @@ -1,7 +1,7 @@ from __future__ import annotations from datetime import UTC, datetime -from typing import Self +from typing import Any, Self from uuid import uuid4 from pydantic import ( diff --git a/generalresearch/models/thl/contest/contest.py b/generalresearch/models/thl/contest/contest.py index bd0fc04..5814bef 100644 --- a/generalresearch/models/thl/contest/contest.py +++ b/generalresearch/models/thl/contest/contest.py @@ -186,7 +186,7 @@ class Contest(ContestBase): @classmethod def model_validate_mysql(cls, data: dict[str, Any]) -> Self: - data = {k: v for k, v in data.items() if k in cls.model_fields.keys()} + data = {k: v for k, v in data.items() if k in cls.model_fields} if isinstance(data["end_condition"], dict): data["end_condition"] = ContestEndCondition.model_validate( data["end_condition"] diff --git a/generalresearch/models/thl/contest/contest_entry.py b/generalresearch/models/thl/contest/contest_entry.py index bb3aef4..b5f0ac3 100644 --- a/generalresearch/models/thl/contest/contest_entry.py +++ b/generalresearch/models/thl/contest/contest_entry.py @@ -1,6 +1,7 @@ from __future__ import annotations from datetime import UTC, datetime +from typing import Any from uuid import uuid4 from pydantic import ( diff --git a/generalresearch/models/thl/contest/raffle.py b/generalresearch/models/thl/contest/raffle.py index 16a0a47..072f011 100644 --- a/generalresearch/models/thl/contest/raffle.py +++ b/generalresearch/models/thl/contest/raffle.py @@ -203,9 +203,7 @@ class RaffleContest(RaffleContestCreate, Contest): c = self.end_condition if c.target_entry_amount and self.current_amount >= c.target_entry_amount: return True - if c.ends_at and datetime.now(tz=UTC) >= c.ends_at: - return True - return False + return bool(c.ends_at and datetime.now(tz=UTC) >= c.ends_at) def model_dump_mysql(self) -> dict[str, Any]: d = super().model_dump_mysql() @@ -213,7 +211,7 @@ class RaffleContest(RaffleContestCreate, Contest): return d @classmethod - def model_validate_mysql(cls, data: dict) -> Self: + def model_validate_mysql(cls, data: dict[str, Any]) -> Self: data["entry_rule"] = ContestEntryRule.model_validate(data["entry_rule"]) return super().model_validate_mysql(data) diff --git a/generalresearch/models/thl/demographics.py b/generalresearch/models/thl/demographics.py index b6a8be1..c11f8b2 100644 --- a/generalresearch/models/thl/demographics.py +++ b/generalresearch/models/thl/demographics.py @@ -76,7 +76,7 @@ class AgeGroup(Enum): return self.label -def calculate_demographic_metrics(opps: list[MarketplaceTask]) -> list: +def calculate_demographic_metrics(opps: list[MarketplaceTask]) -> list[dict[str, Any]]: """ Measurement: marketplace_survey_demographics tags: source (marketplace) @@ -86,7 +86,7 @@ def calculate_demographic_metrics(opps: list[MarketplaceTask]) -> list: """ source = {opp.source for opp in opps} assert len(source) == 1 - source = list(source)[0] + source = next(iter(source)) survey_cpi = defaultdict(list) target_open = defaultdict(int) for opp in opps: @@ -100,7 +100,7 @@ def calculate_demographic_metrics(opps: list[MarketplaceTask]) -> list: survey_counter = {k: len(v) for k, v in survey_cpi.items()} survey_counter = {k: {"count": v} for k, v in survey_counter.items() if v} - grp_stats = dict() + grp_stats = {} for grp, costs in survey_cpi.items(): stats = { "cost_min": np.min(costs), @@ -155,7 +155,7 @@ def calculate_used_question_metrics( """ source = {opp.source for opp in opps} assert len(source) == 1 - source = list(source)[0] + source = next(iter(source)) country_q_counter = defaultdict(Counter) for opp in opps: for q in opp.used_question_ids: diff --git a/generalresearch/models/thl/finance.py b/generalresearch/models/thl/finance.py index 8c94390..79a74a7 100644 --- a/generalresearch/models/thl/finance.py +++ b/generalresearch/models/thl/finance.py @@ -557,7 +557,7 @@ class BusinessBalances(BaseModel): they all explicitly are set """ - if any([pb.product_id is None for pb in v]): + if any(pb.product_id is None for pb in v): raise ValueError("'product_id' must be set for BusinessBalance children.") return v diff --git a/generalresearch/models/thl/ledger.py b/generalresearch/models/thl/ledger.py index 19dde20..a9fbbb1 100644 --- a/generalresearch/models/thl/ledger.py +++ b/generalresearch/models/thl/ledger.py @@ -1,7 +1,7 @@ from __future__ import annotations from datetime import UTC, datetime -from enum import StrEnum +from enum import IntEnum, StrEnum from typing import Annotated, Any, Literal, Self from uuid import uuid4 @@ -36,7 +36,7 @@ from generalresearch.models.thl.payout_format import ( from generalresearch.utils.enum import ReprEnumMeta -class Direction(int, Enum, metaclass=ReprEnumMeta): +class Direction(IntEnum, metaclass=ReprEnumMeta): """Entries on the debit side will increase debit normal accounts, while entries on the credit side will decrease them. Conversely, entries on the credit side will increase credit normal accounts, while entries on diff --git a/generalresearch/models/thl/offerwall/__init__.py b/generalresearch/models/thl/offerwall/__init__.py index e7c8e03..0c3d51d 100644 --- a/generalresearch/models/thl/offerwall/__init__.py +++ b/generalresearch/models/thl/offerwall/__init__.py @@ -267,7 +267,7 @@ class OfferWallRequest(BaseModel): # We need this so thl-core can refresh an offerwall in order to continue # a session d = self.model_dump(mode="json") - kwargs = dict() + kwargs = {} keys = [ "n_bins", "min_bin_size", diff --git a/generalresearch/models/thl/offerwall/base.py b/generalresearch/models/thl/offerwall/base.py index 33489df..33b9847 100644 --- a/generalresearch/models/thl/offerwall/base.py +++ b/generalresearch/models/thl/offerwall/base.py @@ -398,8 +398,10 @@ class OfferwallBucket(BaseModel): ) uri: HttpsUrl | None = Field( examples=[ - "https://task.generalresearch.com/api/v1/52d3f63b2709/797df4136c604a6c8599818296aae6d1/?i" - "=5ba2fe5010cc4d078fc3cc0b0cc264c3&b=test&66482fb=e7baf5e" + ( + "https://task.generalresearch.com/api/v1/52d3f63b2709/797df4136c604a6c8599818296aae6d1/?i" + "=5ba2fe5010cc4d078fc3cc0b0cc264c3&b=test&66482fb=e7baf5e" + ) ], description="The URL to send a respondent into. Must not edit this URL in any way", default=None, diff --git a/generalresearch/models/thl/payout_format.py b/generalresearch/models/thl/payout_format.py index 4d616b6..d29c9de 100644 --- a/generalresearch/models/thl/payout_format.py +++ b/generalresearch/models/thl/payout_format.py @@ -70,7 +70,7 @@ def format_payout_format(payout_format: str, payout_int: int) -> str: except TypeError: # "{payout()*1:}" - TypeError: 'int' object is not callable raise ValueError("Invalid type reference.") - except Exception: + except Exception: # noqa raise ValueError("Invalid payout transformation") formatstr = f"{{:{formatstr}}}" diff --git a/generalresearch/models/thl/product.py b/generalresearch/models/thl/product.py index a7ecd55..65ed177 100644 --- a/generalresearch/models/thl/product.py +++ b/generalresearch/models/thl/product.py @@ -957,7 +957,7 @@ class Product(BaseModel, validate_assignment=True): @field_validator("harmonizer_domain", mode="before") def harmonizer_domain_https(cls, s: str | None): # in the db, this has no scheme. accept both with a default of https:// - if s is not None and not (s.startswith("https://") or s.startswith("http://")): + if s is not None and not (s.startswith(("https://", "http://"))): s = f"https://{s}" return s @@ -1371,7 +1371,7 @@ class Product(BaseModel, validate_assignment=True): if self.payout_config.payout_transformation is None: return None payout_xform_func = self.get_payout_transformation_func() - kwargs = dict() + kwargs = {} if "user_wallet_balance" in inspect.signature(payout_xform_func).parameters: kwargs["user_wallet_balance"] = user_wallet_balance user_payout: Decimal = payout_xform_func(bp_payout, **kwargs) diff --git a/generalresearch/models/thl/profiling/marketplace.py b/generalresearch/models/thl/profiling/marketplace.py index ad4ce80..0129e38 100644 --- a/generalresearch/models/thl/profiling/marketplace.py +++ b/generalresearch/models/thl/profiling/marketplace.py @@ -82,10 +82,9 @@ class MarketplaceQuestion(BaseModel, ABC): # question has more than 6. repr_args = list(self.__repr_args__()) for n, (k, v) in enumerate(repr_args): - if k == "options": - if v and len(v) > 6: - v = v[:3] + ["..."] + v[-3:] - repr_args[n] = ("options", v) + if k == "options" and v and len(v) > 6: + v = v[:3] + ["..."] + v[-3:] + repr_args[n] = ("options", v) join_str = ", " repr_str = join_str.join( repr(v) if a is None else f"{a}={v!r}" for a, v in repr_args diff --git a/generalresearch/models/thl/report_task.py b/generalresearch/models/thl/report_task.py index d29599d..299ba90 100644 --- a/generalresearch/models/thl/report_task.py +++ b/generalresearch/models/thl/report_task.py @@ -28,7 +28,7 @@ def prioritize_report_values( return None report_values = list(set(report_values)) random.shuffle(report_values) - return sorted(report_values, key=lambda x: REPORT_PRIORITY[x])[-1] + return max(report_values, key=lambda x: REPORT_PRIORITY[x]) class ReportTask(BaseModel): diff --git a/generalresearch/models/thl/session.py b/generalresearch/models/thl/session.py index fe7194a..e4e264f 100644 --- a/generalresearch/models/thl/session.py +++ b/generalresearch/models/thl/session.py @@ -234,10 +234,7 @@ class WallBase(BaseModel): return self.is_visible() and self.status == Status.COMPLETE def allow_session(self) -> bool: - if self.status == Status.COMPLETE: - return False - - return True + return self.status != Status.COMPLETE def update(self, **kwargs) -> None: """ @@ -969,10 +966,7 @@ class Session(BaseModel): return True # Hard limit of 40 wall events per session - if len(self.wall_events) >= 40: - return True - - return False + return len(self.wall_events) >= 40 def determine_payments( self, @@ -985,6 +979,7 @@ class Session(BaseModel): ) product = self.user.product + assert product # Handle brokerage product payouts bp_pay: Decimal = product.determine_bp_payment(thl_net) commission_amount: Decimal = thl_net - bp_pay diff --git a/generalresearch/models/thl/soft_pair.py b/generalresearch/models/thl/soft_pair.py index f3b2b6f..7c2f36e 100644 --- a/generalresearch/models/thl/soft_pair.py +++ b/generalresearch/models/thl/soft_pair.py @@ -50,7 +50,7 @@ class SoftPairResult: return ( self.survey_id + ":" - + ";".join(sorted(set([c.question_id for c in self.conditions]))) + + ";".join(sorted({c.question_id for c in self.conditions})) ) else: return None diff --git a/generalresearch/models/thl/survey/penalty.py b/generalresearch/models/thl/survey/penalty.py index 04f8e20..755d25c 100644 --- a/generalresearch/models/thl/survey/penalty.py +++ b/generalresearch/models/thl/survey/penalty.py @@ -56,7 +56,7 @@ class TeamSurveyPenalty(SurveyPenalty): Penalty = Annotated[ - Union[BPSurveyPenalty, TeamSurveyPenalty], + BPSurveyPenalty | TeamSurveyPenalty, Field(discriminator="kind"), ] PenaltyListAdapter = TypeAdapter(list[Penalty]) diff --git a/generalresearch/models/thl/survey/task_collection.py b/generalresearch/models/thl/survey/task_collection.py index b80166f..d8db0d1 100644 --- a/generalresearch/models/thl/survey/task_collection.py +++ b/generalresearch/models/thl/survey/task_collection.py @@ -38,7 +38,8 @@ class TaskCollection(BaseModel): except pa.errors.SchemaErrors as exc: idx = exc.failure_cases["index"] if len(idx) >= len(df) * 0.10: - raise exc + raise + logger.info(f"{self.__repr_name__()}:handle_df:{json.dumps(exc.message)}") df.drop(index=list(idx), inplace=True) # we need to redo the validation after removing failing rows! diff --git a/generalresearch/models/thl/task_status.py b/generalresearch/models/thl/task_status.py index de767d6..7719b18 100644 --- a/generalresearch/models/thl/task_status.py +++ b/generalresearch/models/thl/task_status.py @@ -224,11 +224,12 @@ class TaskStatusResponse(BaseModel): return v or 0 @field_validator("kwargs", mode="after") - def sanitize_kwargs(cls, v: dict | None) -> dict | None: + def sanitize_kwargs(cls, v: dict[str, Any] | None) -> dict[str, Any] | None: if v and "clicked_timestamp" in v: try: - clicked_timestamp = datetime.strptime( - v["clicked_timestamp"], "%Y-%m-%d %H:%M:%S.%f" + clicked_timestamp = datetime.strptime( # noqa + date_string=v["clicked_timestamp"], + format="%Y-%m-%d %H:%M:%S.%f", ) v["clicked_timestamp"] = ( clicked_timestamp.isoformat(timespec="microseconds") + "Z" @@ -238,7 +239,7 @@ class TaskStatusResponse(BaseModel): return v @model_validator(mode="before") - def transform_user_payout(cls, d): + def transform_user_payout(cls, d: dict[str, Any]): # If the user_payout is None and there is a payout_format, make the user_payout 0 if d.get("user_payout") is None and d.get("payout_format"): d["user_payout"] = 0 diff --git a/generalresearch/pg_helper.py b/generalresearch/pg_helper.py index 1d5d30b..a397247 100644 --- a/generalresearch/pg_helper.py +++ b/generalresearch/pg_helper.py @@ -108,10 +108,8 @@ class PostgresConfig: def execute_write(self, query, params=None) -> int: cmd = query.lstrip().upper() - assert ( - cmd.startswith("INSERT") - or cmd.startswith("UPDATE") - or cmd.startswith("DELETE") + assert cmd.startswith( + ("INSERT", "UPDATE", "DELETE") ), "Supports INSERT/UPDATE only" with self.make_connection() as conn: diff --git a/generalresearch/sql_helper.py b/generalresearch/sql_helper.py index 08b660d..ae2b8d8 100644 --- a/generalresearch/sql_helper.py +++ b/generalresearch/sql_helper.py @@ -315,7 +315,7 @@ class SqlHelper(SqlConnector): field_names = ["`" + x + "`" for x in field_names] field_name_str = ",".join(field_names) if filter_d: - lookup_vals = " AND ".join([f"`{fn}`=%({fn})s" for fn in filter_d.keys()]) + lookup_vals = " AND ".join([f"`{fn}`=%({fn})s" for fn in filter_d]) lookup_str = f" WHERE {lookup_vals}" else: lookup_str = "" diff --git a/generalresearch/utils/grpc_logger.py b/generalresearch/utils/grpc_logger.py index 59f7471..8f2f454 100644 --- a/generalresearch/utils/grpc_logger.py +++ b/generalresearch/utils/grpc_logger.py @@ -33,9 +33,11 @@ try: response = handler_func(request, context) code = context.code() or grpc.StatusCode.OK return response - except Exception as e: + + except Exception: code = context.code() or grpc.StatusCode.INTERNAL - raise e + raise + finally: duration_ms = int((time.time() - start_time) * 1000) peer = context.peer() or "unknown" diff --git a/generalresearch/wall_status_codes/fullcircle.py b/generalresearch/wall_status_codes/fullcircle.py index aeaa4c7..cd9fdff 100644 --- a/generalresearch/wall_status_codes/fullcircle.py +++ b/generalresearch/wall_status_codes/fullcircle.py @@ -29,7 +29,7 @@ status_codes_ext_map: dict[StatusCode1, list[str]] = { StatusCode1.PS_FAIL: [], StatusCode1.PS_OVERQUOTA: [], } -ext_status_code_map: dict[str, StatusCode1] = dict() +ext_status_code_map: dict[str, StatusCode1] = {} for k, v in status_codes_ext_map.items(): k: StatusCode1 v: list[str] diff --git a/generalresearch/wall_status_codes/innovate.py b/generalresearch/wall_status_codes/innovate.py index 936ee6c..e3d2468 100644 --- a/generalresearch/wall_status_codes/innovate.py +++ b/generalresearch/wall_status_codes/innovate.py @@ -38,7 +38,7 @@ status_codes_ext_map: dict[StatusCode1, list[str]] = { StatusCode1.PS_FAIL: ["5"], StatusCode1.PS_OVERQUOTA: ["7"], } -ext_status_code_map = dict() +ext_status_code_map = {} for k, v in status_codes_ext_map.items(): for vv in v: ext_status_code_map[status_codes_ext_map.get(vv, vv)] = k diff --git a/generalresearch/wall_status_codes/lucid.py b/generalresearch/wall_status_codes/lucid.py index 3cc1b5e..c4c5e90 100644 --- a/generalresearch/wall_status_codes/lucid.py +++ b/generalresearch/wall_status_codes/lucid.py @@ -102,7 +102,7 @@ status_codes_ext_map: dict[StatusCode1, list[str]] = { StatusCode1.PS_OVERQUOTA: ["40", "41", "42"], } -ext_status_code_map: dict[str, StatusCode1] = dict() +ext_status_code_map: dict[str, StatusCode1] = {} for k, v in status_codes_ext_map.items(): k: StatusCode1 v: list[str] diff --git a/generalresearch/wall_status_codes/morning.py b/generalresearch/wall_status_codes/morning.py index ffa6be2..6f63b82 100644 --- a/generalresearch/wall_status_codes/morning.py +++ b/generalresearch/wall_status_codes/morning.py @@ -97,7 +97,7 @@ status_codes_ext_map: dict[StatusCode1, list[str]] = { "quota_invalid_for_bid", ], } -ext_status_code_map: dict[str, StatusCode1] = dict() +ext_status_code_map: dict[str, StatusCode1] = {} for k, v in status_codes_ext_map.items(): k: StatusCode1 v: list[str] diff --git a/generalresearch/wall_status_codes/pollfish.py b/generalresearch/wall_status_codes/pollfish.py index a5c6e25..e1ad12a 100644 --- a/generalresearch/wall_status_codes/pollfish.py +++ b/generalresearch/wall_status_codes/pollfish.py @@ -58,7 +58,7 @@ status_codes_ext_map: dict[StatusCode1, list[str]] = { ], StatusCode1.PS_OVERQUOTA: ["quota_full", "survey_closed", "survey_expired"], } -ext_status_code_map = dict() +ext_status_code_map = {} for k, v in status_codes_ext_map.items(): k: StatusCode1 v: list[str] diff --git a/test_utils/managers/contest/conftest.py b/test_utils/managers/contest/conftest.py index 67935e7..a9375f6 100644 --- a/test_utils/managers/contest/conftest.py +++ b/test_utils/managers/contest/conftest.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import pytest from generalresearch.managers.base import Permission @@ -11,8 +13,6 @@ def contest_manager(thl_web_rw: PostgresConfig) -> ContestManager: assert thl_web_rw.dsn.path assert "/unittest-" in thl_web_rw.dsn.path - from generalresearch.managers.thl.contest_manager import ContestManager - return ContestManager( pg_config=thl_web_rw, permissions=[ diff --git a/test_utils/models/contest/conftest.py b/test_utils/models/contest/conftest.py index d9c8a6b..84930b8 100644 --- a/test_utils/models/contest/conftest.py +++ b/test_utils/models/contest/conftest.py @@ -38,6 +38,26 @@ from generalresearch.models.thl.user import User # === Managers === +# --- Factories --- + + +@pytest.fixture(scope="function") +def raffle_contest_factory( + product_user_wallet_yes: Product, + raffle_contest_create: RaffleContestCreate, + contest_manager: ContestManager, +) -> Callable[..., RaffleContest]: + + def _inner(**kwargs): + raffle_contest_create.update(**kwargs) + return contest_manager.create( + product_id=product_user_wallet_yes.uuid, + contest_create=raffle_contest_create, + ) + + return _inner + + # === Models === @@ -82,23 +102,6 @@ def raffle_contest( ) -@pytest.fixture(scope="function") -def raffle_contest_factory( - product_user_wallet_yes: Product, - raffle_contest_create: RaffleContestCreate, - contest_manager: ContestManager, -) -> Callable[..., RaffleContest]: - - def _inner(**kwargs): - raffle_contest_create.update(**kwargs) - return contest_manager.create( - product_id=product_user_wallet_yes.uuid, - contest_create=raffle_contest_create, - ) - - return _inner - - @pytest.fixture def milestone_contest_create() -> MilestoneContestCreate: from generalresearch.models.thl.contest import ( diff --git a/test_utils/spectrum/conftest.py b/test_utils/spectrum/conftest.py index 7cd9321..a8ce9d9 100644 --- a/test_utils/spectrum/conftest.py +++ b/test_utils/spectrum/conftest.py @@ -1,32 +1,32 @@ from __future__ import annotations -import logging import time from datetime import UTC, datetime from decimal import Decimal -from typing import TYPE_CHECKING, Any +from typing import Any import pytest +from generalresearch.config import GRLBaseSettings from generalresearch.managers.spectrum.survey import ( SpectrumCriteriaManager, SpectrumSurveyManager, ) -from generalresearch.models.spectrum.survey import SpectrumSurvey +from generalresearch.models import ( + LogicalOperator, +) +from generalresearch.models.spectrum.survey import ( + SpectrumCondition, + SpectrumSurvey, +) +from generalresearch.models.thl.survey.condition import ConditionValueType from generalresearch.sql_helper import SqlHelper -from .surveys_json import CONDITIONS, SURVEYS_JSON - -if TYPE_CHECKING: - from generalresearch.config import GRLBaseSettings - @pytest.fixture(scope="session") def spectrum_rw(settings: GRLBaseSettings) -> SqlHelper: - logging.info(f"{settings.spectrum_rw_db=}") - assert settings.spectrum_rw_db is not None - assert "/unittest-" in settings.spectrum_rw_db.path + assert "/unittest-" in str(settings.spectrum_rw_db.path) return SqlHelper( dsn=settings.spectrum_rw_db, @@ -38,27 +38,36 @@ def spectrum_rw(settings: GRLBaseSettings) -> SqlHelper: @pytest.fixture(scope="session") def spectrum_criteria_manager(spectrum_rw: SqlHelper) -> SpectrumCriteriaManager: + assert spectrum_rw.dsn + assert spectrum_rw.dsn.path assert "/unittest-" in spectrum_rw.dsn.path return SpectrumCriteriaManager(spectrum_rw) @pytest.fixture(scope="session") def spectrum_survey_manager(spectrum_rw: SqlHelper) -> SpectrumSurveyManager: + assert spectrum_rw.dsn + assert spectrum_rw.dsn.path assert "/unittest-" in spectrum_rw.dsn.path return SpectrumSurveyManager(spectrum_rw) @pytest.fixture(scope="session") def setup_spectrum_surveys( - spectrum_rw: SqlHelper, spectrum_survey_manager, spectrum_criteria_manager + spectrum_rw: SqlHelper, + spectrum_survey_manager: SpectrumSurveyManager, + spectrum_criteria_manager: SpectrumCriteriaManager, + spectrum_conditions: list[SpectrumCondition], + spectrum_api_surveys_json: list[str], ) -> None: now = datetime.now(UTC) # make sure these example surveys exist in db - surveys = [SpectrumSurvey.model_validate_json(x) for x in SURVEYS_JSON] + surveys = [SpectrumSurvey.model_validate_json(x) for x in spectrum_api_surveys_json] for s in surveys: s.modified_api = datetime.now(tz=UTC) + spectrum_survey_manager.create_or_update(surveys) - spectrum_criteria_manager.update(CONDITIONS) + spectrum_criteria_manager.update(spectrum_conditions) # and make sure they have allocation for 687 spectrum_rw.execute_sql_query( @@ -198,6 +207,46 @@ def spectrum_api_surveys_json() -> list[str]: ] +def spectrum_conditions() -> list[SpectrumCondition]: + # make sure hashes for 111111 are in db + c1 = SpectrumCondition( + question_id="1001", + value_type=ConditionValueType.LIST, + values=["a", "b", "c"], + negate=False, + logical_operator=LogicalOperator.OR, + ) + c2 = SpectrumCondition( + question_id="1001", + value_type=ConditionValueType.LIST, + values=["a"], + negate=False, + logical_operator=LogicalOperator.OR, + ) + c3 = SpectrumCondition( + question_id="1002", + value_type=ConditionValueType.RANGE, + values=["18-24", "30-32"], + negate=False, + logical_operator=LogicalOperator.OR, + ) + c4 = SpectrumCondition( + question_id="212", + value_type=ConditionValueType.LIST, + values=["23", "24"], + negate=False, + logical_operator=LogicalOperator.OR, + ) + c5 = SpectrumCondition( + question_id="1031", + value_type=ConditionValueType.LIST, + values=["113", "114", "121"], + negate=False, + logical_operator=LogicalOperator.OR, + ) + return [c1, c2, c3, c4, c5] + + @pytest.fixture(scope="session") def spectrum_api_survey_json() -> dict[str, Any]: return { diff --git a/tests/incite/collections/test_df_collection_base.py b/tests/incite/collections/test_df_collection_base.py index b9f0181..c236700 100644 --- a/tests/incite/collections/test_df_collection_base.py +++ b/tests/incite/collections/test_df_collection_base.py @@ -9,6 +9,7 @@ from generalresearch.incite.collections import ( DFCollection, DFCollectionType, ) +from generalresearch.pg_helper import PostgresConfig if TYPE_CHECKING: from generalresearch.incite.base import GRLDatasets @@ -45,7 +46,9 @@ class TestDFCollectionBase: class TestDFCollectionBaseProperties: @pytest.mark.skip - def test_df_collection_items(self, mnt_filepath: GRLDatasets, df_coll_type): + def test_df_collection_items( + self, mnt_filepath: GRLDatasets, df_coll_type: DFCollectionType + ): instance = DFCollection( data_type=df_coll_type, start=datetime(year=1800, month=1, day=1, tzinfo=UTC), @@ -57,7 +60,9 @@ class TestDFCollectionBaseProperties: assert len(instance.interval_range) == len(instance.items) assert len(instance.items) == 366 - def test_df_collection_progress(self, mnt_filepath: GRLDatasets, df_coll_type): + def test_df_collection_progress( + self, mnt_filepath: GRLDatasets, df_coll_type: DFCollectionType + ): instance = DFCollection( data_type=df_coll_type, start=datetime(year=1800, month=1, day=1, tzinfo=UTC), @@ -70,7 +75,9 @@ class TestDFCollectionBaseProperties: assert isinstance(instance.progress, pd.DataFrame) assert instance.progress.shape == (366, 6) - def test_df_collection_schema(self, mnt_filepath: GRLDatasets, df_coll_type): + def test_df_collection_schema( + self, mnt_filepath: GRLDatasets, df_coll_type: DFCollectionType + ): instance1 = DFCollection( data_type=DFCollectionType.WALL, archive_path=mnt_filepath.data_src ) @@ -87,9 +94,9 @@ class TestDFCollectionBaseProperties: class TestDFCollectionBaseMethods: @pytest.mark.skip - def test_initial_load(self, mnt_filepath: GRLDatasets, thl_web_rr): + def test_initial_load(self, mnt_filepath: GRLDatasets, thl_web_rr: PostgresConfig): instance = DFCollection( - pg_config=thl_web_rr: PostgresConfig, + pg_config=thl_web_rr, data_type=DFCollectionType.USER, start=datetime(year=2022, month=1, day=1, minute=0, tzinfo=UTC), finished=datetime(year=2022, month=1, day=1, minute=5, tzinfo=UTC), diff --git a/tests/incite/collections/test_df_collection_item_base.py b/tests/incite/collections/test_df_collection_item_base.py index 9a2ecf3..e0171c2 100644 --- a/tests/incite/collections/test_df_collection_item_base.py +++ b/tests/incite/collections/test_df_collection_item_base.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from datetime import UTC, datetime from typing import TYPE_CHECKING @@ -19,7 +21,7 @@ df_collection_types = [e for e in DFCollectionType if e is not DFCollectionType. @pytest.mark.parametrize("df_coll_type", df_collection_types) class TestDFCollectionItemBase: - def test_init(self, mnt_filepath: GRLDatasets, df_coll_type): + def test_init(self, mnt_filepath: GRLDatasets, df_coll_type: DFCollectionType): collection = DFCollection( data_type=df_coll_type, offset="100d", @@ -38,14 +40,16 @@ class TestDFCollectionItemBase: class TestDFCollectionItemProperties: @pytest.mark.skip - def test_filename(self, df_coll_type): + def test_filename(self, df_coll_type: DFCollectionType): pass @pytest.mark.parametrize("df_coll_type", df_collection_types) class TestDFCollectionItemMethods: - def test_has_mysql_false(self, mnt_filepath: GRLDatasets, df_coll_type): + def test_has_mysql_false( + self, mnt_filepath: GRLDatasets, df_coll_type: DFCollectionType + ): collection = DFCollection( data_type=df_coll_type, offset="100d", @@ -58,7 +62,10 @@ class TestDFCollectionItemMethods: assert not instance1.has_mysql() def test_has_mysql_true( - self, thl_web_rr: PostgresConfig, mnt_filepath: GRLDatasets, df_coll_type + self, + thl_web_rr: PostgresConfig, + mnt_filepath: GRLDatasets, + df_coll_type: DFCollectionType, ): collection = DFCollection( data_type=df_coll_type, @@ -66,7 +73,7 @@ class TestDFCollectionItemMethods: start=datetime(year=1800, month=1, day=1, tzinfo=UTC), finished=datetime(year=1900, month=1, day=1, tzinfo=UTC), archive_path=mnt_filepath.archive_path(enum_type=df_coll_type), - pg_config=thl_web_rr: PostgresConfig, + pg_config=thl_web_rr, ) # Has RR, assume unittest server is online @@ -74,5 +81,5 @@ class TestDFCollectionItemMethods: assert instance2.has_mysql() @pytest.mark.skip - def test_update_partial_archive(self, df_coll_type): + def test_update_partial_archive(self, df_coll_type: DFCollectionType): pass diff --git a/tests/incite/collections/test_df_collection_thl_marketplaces.py b/tests/incite/collections/test_df_collection_thl_marketplaces.py index d2d3ce4..b4b5b00 100644 --- a/tests/incite/collections/test_df_collection_thl_marketplaces.py +++ b/tests/incite/collections/test_df_collection_thl_marketplaces.py @@ -4,6 +4,7 @@ from itertools import product import pytest from pandera.pandas import Column, DataFrameSchema, Index +from generalresearch.incite.base import GRLDatasets from generalresearch.incite.collections import DFCollection, DFCollectionType from generalresearch.incite.collections.thl_marketplaces import ( InnovateSurveyHistoryCollection, @@ -11,6 +12,7 @@ from generalresearch.incite.collections.thl_marketplaces import ( SagoSurveyHistoryCollection, SpectrumSurveyTimeseriesCollection, ) +from generalresearch.pg_helper import PostgresConfig def combo_object(): @@ -29,7 +31,13 @@ def combo_object(): @pytest.mark.parametrize("df_coll, offset", combo_object()) class TestDFCollection_thl_marketplaces: - def test_init(self, mnt_filepath, df_coll, offset, spectrum_rw): + def test_init( + self, + mnt_filepath: GRLDatasets, + df_coll: DFCollection, + offset: str, + spectrum_rw: PostgresConfig, + ): assert issubclass(df_coll, DFCollection) # This is stupid, but we need to pull the default from the @@ -38,7 +46,7 @@ class TestDFCollection_thl_marketplaces: assert isinstance(data_type, DFCollectionType) # (1) Can't be totally empty, needs a path... - with pytest.raises(expected_exception=Exception) as cm: + with pytest.raises(expected_exception=Exception): instance = df_coll() # (2) Confirm it only needs the archive_path @@ -61,7 +69,7 @@ class TestDFCollection_thl_marketplaces: assert isinstance(instance._schema, DataFrameSchema) assert isinstance(instance._schema.index, Index) - for c in instance._schema.columns.keys(): + for c in instance._schema.columns: assert isinstance(c, str) col = instance._schema.columns[c] assert isinstance(col, Column) diff --git a/tests/incite/collections/test_df_collection_thl_web.py b/tests/incite/collections/test_df_collection_thl_web.py index bcdeb83..6d509bc 100644 --- a/tests/incite/collections/test_df_collection_thl_web.py +++ b/tests/incite/collections/test_df_collection_thl_web.py @@ -3,19 +3,16 @@ from __future__ import annotations from collections.abc import Generator from datetime import datetime from itertools import product -from typing import TYPE_CHECKING import dask.dataframe as dd import pandas as pd import pytest from pandera.pandas import DataFrameSchema -from generalresearch.incite.collections import DFCollection, DFCollectionType - -if TYPE_CHECKING: - from generalresearch.incite.collections import ( - DFCollectionType, - ) +from generalresearch.incite.collections import ( + DFCollection, + DFCollectionType, +) def combo_object() -> Generator[tuple]: @@ -39,7 +36,10 @@ def combo_object() -> Generator[tuple]: class TestDFCollection_thl_web: def test_init( - self, df_collection_data_type: DFCollectionType, offset: str, df_collection + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, ): assert isinstance(df_collection_data_type, DFCollectionType) assert isinstance(df_collection, DFCollection) @@ -50,12 +50,12 @@ class TestDFCollection_thl_web: ) class TestDFCollection_thl_web_Properties: - def test_items(self, df_collection): + def test_items(self, df_collection: DFCollection): assert isinstance(df_collection.items, list) for i in df_collection.items: assert i._collection == df_collection - def test__schema(self, df_collection): + def test__schema(self, df_collection: DFCollection): assert isinstance(df_collection._schema, DataFrameSchema) @@ -65,16 +65,16 @@ class TestDFCollection_thl_web_Properties: class TestDFCollection_thl_web_BaseProperties: @pytest.mark.skip - def test__interval_range(self, df_collection): + def test__interval_range(self, df_collection: DFCollection): pass - def test_interval_start(self, df_collection): + def test_interval_start(self, df_collection: DFCollection): assert isinstance(df_collection.interval_start, datetime) - def test_interval_range(self, df_collection): + def test_interval_range(self, df_collection: DFCollection): assert isinstance(df_collection.interval_range, list) - def test_progress(self, df_collection): + def test_progress(self, df_collection: DFCollection): assert isinstance(df_collection.progress, pd.DataFrame) @@ -84,17 +84,21 @@ class TestDFCollection_thl_web_BaseProperties: class TestDFCollection_thl_web_Methods: @pytest.mark.skip - def test_initial_loads(self, df_collection_data_type, df_collection, offset): + def test_initial_loads( + self, df_collection_data_type, df_collection: DFCollection, offset: str + ): pass @pytest.mark.skip def test_fetch_force_rr_latest( - self, df_collection_data_type, df_collection, offset: str + self, df_collection_data_type, df_collection: DFCollection, offset: str ): pass @pytest.mark.skip - def test_force_rr_latest(self, df_collection_data_type, df_collection, offset): + def test_force_rr_latest( + self, df_collection_data_type, df_collection: DFCollection, offset: str + ): pass @@ -103,63 +107,108 @@ class TestDFCollection_thl_web_Methods: ) class TestDFCollection_thl_web_BaseMethods: - def test_fetch_all_paths(self, df_collection_data_type, offset: str, df_collection): + def test_fetch_all_paths( + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, + ): res = df_collection.fetch_all_paths( items=None, force_rr_latest=False, include_partial=False ) assert isinstance(res, list) @pytest.mark.skip - def test_ddf(self, df_collection_data_type, offset: str, df_collection): + def test_ddf( + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, + ): res = df_collection.ddf() assert isinstance(res, dd.DataFrame) # -- cleanup -- @pytest.mark.skip def test_schedule_cleanup( - self, df_collection_data_type, offset: str, df_collection + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, ): pass @pytest.mark.skip - def test_cleanup(self, df_collection_data_type, offset: str, df_collection): + def test_cleanup( + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, + ): pass @pytest.mark.skip def test_cleanup_partials( - self, df_collection_data_type, offset: str, df_collection + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, ): pass @pytest.mark.skip def test_clear_tmp_archives( - self, df_collection_data_type, offset: str, df_collection + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, ): pass @pytest.mark.skip def test_clear_corrupt_archives( - self, df_collection_data_type, offset: str, df_collection + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, ): pass @pytest.mark.skip def test_rebuild_symlinks( - self, df_collection_data_type, offset: str, df_collection + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, ): pass # -- Source timing -- @pytest.mark.skip - def test_get_item(self, df_collection_data_type, offset: str, df_collection): + def test_get_item( + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, + ): pass @pytest.mark.skip - def test_get_item_start(self, df_collection_data_type, offset: str, df_collection): + def test_get_item_start( + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, + ): pass @pytest.mark.skip - def test_get_items(self, df_collection_data_type, offset: str, df_collection): + def test_get_items( + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, + ): # If we get all the items from the start of the collection, it # should include all the items! res1 = df_collection.items @@ -168,18 +217,27 @@ class TestDFCollection_thl_web_BaseMethods: @pytest.mark.skip def test_get_items_from_year( - self, df_collection_data_type, offset: str, df_collection + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, ): pass @pytest.mark.skip def test_get_items_last90( - self, df_collection_data_type, offset: str, df_collection + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, ): pass @pytest.mark.skip def test_get_items_last365( - self, df_collection_data_type, offset: str, df_collection + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, ): pass diff --git a/tests/incite/mergers/foundations/test_user_id_product.py b/tests/incite/mergers/foundations/test_user_id_product.py index 10802e5..7367056 100644 --- a/tests/incite/mergers/foundations/test_user_id_product.py +++ b/tests/incite/mergers/foundations/test_user_id_product.py @@ -1,11 +1,15 @@ +from __future__ import annotations + from datetime import UTC, datetime, timedelta from itertools import product import pandas as pd import pytest +from dask.distributed import Client as DaskClient # noinspection PyUnresolvedReferences from generalresearch.incite.mergers.foundations.user_id_product import ( + UserIdProductMerge, UserIdProductMergeItem, ) @@ -23,14 +27,21 @@ from generalresearch.incite.mergers.foundations.user_id_product import ( class TestUserIDProduct: @pytest.mark.skip - def test_base(self, client_no_amm, user_id_product_merge): + def test_base( + self, client_no_amm: DaskClient, user_id_product_merge: UserIdProductMerge + ): ddf = user_id_product_merge.ddf() df = client_no_amm.compute(collections=ddf, sync=True) assert isinstance(df, pd.DataFrame) assert not df.empty @pytest.mark.skip - def test_base_item(self, client_no_amm, user_id_product_merge, user_collection): + def test_base_item( + self, + client_no_amm: DaskClient, + user_id_product_merge: UserIdProductMerge, + user_collection, + ): assert len(user_id_product_merge.items) == 1 for item in user_id_product_merge.items: @@ -40,7 +51,7 @@ class TestUserIDProduct: try: modified_time1 = path.stat().st_mtime - except Exception: + except OSError: modified_time1 = 0 user_id_product_merge.build(client=client_no_amm, user_coll=user_collection) @@ -49,7 +60,9 @@ class TestUserIDProduct: assert modified_time2 > modified_time1 @pytest.mark.skip - def test_read(self, client_no_amm, user_id_product_merge): + def test_read( + self, client_no_amm: DaskClient, user_id_product_merge: UserIdProductMerge + ): users_ddf = user_id_product_merge.ddf() df = client_no_amm.compute(collections=users_ddf, sync=True) diff --git a/tests/incite/mergers/test_pop_ledger.py b/tests/incite/mergers/test_pop_ledger.py index 529a641..2146344 100644 --- a/tests/incite/mergers/test_pop_ledger.py +++ b/tests/incite/mergers/test_pop_ledger.py @@ -86,9 +86,7 @@ class TestMergePOPLedger: # -- - user_wallet_account: LedgerAccount = ( - thl_ledger_manager.get_account_or_create_user_wallet(user=u) - ) + thl_ledger_manager.get_account_or_create_user_wallet(user=u) cash_account: LedgerAccount = thl_ledger_manager.get_account_cash() rev_account: LedgerAccount = ( thl_ledger_manager.get_account_task_complete_revenue() @@ -295,7 +293,7 @@ class TestMergePOPLedger: assert isinstance(df.index, pd.Index) assert isinstance(df.index, pd.DatetimeIndex) - bp_account_balance = thl_ledger_manager.get_account_balance(account=bp_account) + thl_ledger_manager.get_account_balance(account=bp_account) # Initial sum initial_sum = df.sum().sum() diff --git a/tests/incite/test_collection_base.py b/tests/incite/test_collection_base.py index d6ce2b1..577eda9 100644 --- a/tests/incite/test_collection_base.py +++ b/tests/incite/test_collection_base.py @@ -241,7 +241,7 @@ class TestCollectionBaseMethodsCleanup: assert "Must override" in str(cm.value) -class TestCollectionBaseMethodsCleanup: +class TestCollectionBaseMethodsCleanup2: @pytest.mark.skip def test_cleanup_partials(self, mnt_filepath: GRLDatasets): diff --git a/tests/managers/thl/test_contest/test_milestone.py b/tests/managers/thl/test_contest/test_milestone.py index e3889bc..aa738e1 100644 --- a/tests/managers/thl/test_contest/test_milestone.py +++ b/tests/managers/thl/test_contest/test_milestone.py @@ -15,6 +15,7 @@ from generalresearch.models.thl.contest.milestone import ( MilestoneContestCreate, MilestoneUserView, ) +from generalresearch.models.thl.contest.raffle import RaffleContest from generalresearch.models.thl.product import Product from generalresearch.models.thl.user import User @@ -241,7 +242,7 @@ class TestMilestoneContestUserViews: def test_list_user_eligible_country( self, user_with_wallet: User, - contest_factory: Callable[..., Contest], + raffle_contest_factory: Callable[..., Contest], thl_ledger_manager: ThlLedgerManager, contest_manager: ContestManager, ): @@ -252,7 +253,7 @@ class TestMilestoneContestUserViews: assert len(cs) == 0 # Create a contest. It'll be in the US/CA - contest_factory(country_isos={"us", "ca"}) + raffle_contest_factory(country_isos={"us", "ca"}) # Not eligible in mexico cs = contest_manager.get_many_by_user_eligible( @@ -265,7 +266,7 @@ class TestMilestoneContestUserViews: assert len(cs) == 1 # Create another, any country - contest_factory(country_isos=None) + raffle_contest_factory(country_isos=None) cs = contest_manager.get_many_by_user_eligible( user=user_with_wallet, country_iso="mx" ) @@ -278,12 +279,12 @@ class TestMilestoneContestUserViews: def test_list_user_eligible( self, user_with_money: User, - contest_factory: Callable[..., Contest], + raffle_contest_factory: Callable[..., RaffleContest], thl_ledger_manager: ThlLedgerManager, contest_manager: ContestManager, ): # User reaches milestone after 1 complete - c = contest_factory(target_amount=1) + c = raffle_contest_factory(target_amount=1) user = user_with_money cs = contest_manager.get_many_by_user_eligible( diff --git a/tests/managers/thl/test_ledger/test_thl_lm_tx__user_payouts.py b/tests/managers/thl/test_ledger/test_thl_lm_tx__user_payouts.py index 5fb6935..82dc143 100644 --- a/tests/managers/thl/test_ledger/test_thl_lm_tx__user_payouts.py +++ b/tests/managers/thl/test_ledger/test_thl_lm_tx__user_payouts.py @@ -277,11 +277,12 @@ class TestLedgerManagerAMT: thl_ledger_manager.create_tx_user_payout_cancelled( user, payout_event=pe, skip_flag_check=True ) - with pytest.raises(expected_exception=LedgerTransactionConditionFailedError): - with caplog.at_level(logging.WARNING): - thl_ledger_manager.create_tx_user_payout_complete( - user, payout_event=pe, skip_flag_check=True - ) + with pytest.raises( + expected_exception=LedgerTransactionConditionFailedError + ), caplog.at_level(logging.WARNING): + thl_ledger_manager.create_tx_user_payout_complete( + user, payout_event=pe, skip_flag_check=True + ) assert "trying to complete payout that was already cancelled" in caplog.text cash = thl_ledger_manager.get_account_cash() @@ -319,11 +320,12 @@ class TestLedgerManagerAMT: thl_ledger_manager.create_tx_user_payout_complete( user, payout_event=pe2, skip_flag_check=True ) - with pytest.raises(expected_exception=LedgerTransactionConditionFailedError): - with caplog.at_level(logging.WARNING): - thl_ledger_manager.create_tx_user_payout_cancelled( - user, payout_event=pe2, skip_flag_check=True - ) + with pytest.raises( + expected_exception=LedgerTransactionConditionFailedError + ), caplog.at_level(logging.WARNING): + thl_ledger_manager.create_tx_user_payout_cancelled( + user, payout_event=pe2, skip_flag_check=True + ) assert "trying to cancel payout that was already completed" in caplog.text diff --git a/tests/models/spectrum/test_survey.py b/tests/models/spectrum/test_survey.py index f97860f..bad6857 100644 --- a/tests/models/spectrum/test_survey.py +++ b/tests/models/spectrum/test_survey.py @@ -407,44 +407,12 @@ class TestSpectrumSurvey: ) -def test_spectrum_something(spectrum_api_surveys_json: list[str]): - # make sure hashes for 111111 are in db - c1 = SpectrumCondition( - question_id="1001", - value_type=ConditionValueType.LIST, - values=["a", "b", "c"], - negate=False, - logical_operator=LogicalOperator.OR, - ) - c2 = SpectrumCondition( - question_id="1001", - value_type=ConditionValueType.LIST, - values=["a"], - negate=False, - logical_operator=LogicalOperator.OR, - ) - c3 = SpectrumCondition( - question_id="1002", - value_type=ConditionValueType.RANGE, - values=["18-24", "30-32"], - negate=False, - logical_operator=LogicalOperator.OR, - ) - c4 = SpectrumCondition( - question_id="212", - value_type=ConditionValueType.LIST, - values=["23", "24"], - negate=False, - logical_operator=LogicalOperator.OR, - ) - c5 = SpectrumCondition( - question_id="1031", - value_type=ConditionValueType.LIST, - values=["113", "114", "121"], - negate=False, - logical_operator=LogicalOperator.OR, - ) - _conditions = [c1, c2, c3, c4, c5] +def test_spectrum_something( + spectrum_conditions: list[SpectrumCondition], spectrum_api_surveys_json: list[str] +): + + c1 = spectrum_conditions[0] + c3 = spectrum_conditions[2] survey = SpectrumSurvey.model_validate_json(spectrum_api_surveys_json[0]) assert c1.criterion_hash in survey.qualifications -- cgit v1.2.3 From 89e88d7695044785470824dee839fd3cd6bd5ddf Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Sun, 30 Aug 2026 22:14:44 -0700 Subject: imports work, gr base_init works. jenkins test1 --- Jenkinsfile | 189 +----------------------------- generalresearch/incite/base.py | 10 +- generalresearch/managers/thl/ipinfo.py | 4 +- generalresearch/models/gr/__init__.py | 22 ++-- generalresearch/models/gr/business.py | 4 +- generalresearch/models/thl/session.py | 2 +- test_utils/conftest.py | 42 ++++--- test_utils/managers/cashout_methods.py | 85 ++++++++------ test_utils/managers/conftest.py | 7 +- test_utils/managers/gr/conftest.py | 4 +- test_utils/models/thl/conftest.py | 12 +- tests/managers/thl/test_cashout_method.py | 7 +- tests/models/gr/test_base.py | 9 +- 13 files changed, 119 insertions(+), 278 deletions(-) (limited to 'test_utils/models') diff --git a/Jenkinsfile b/Jenkinsfile index e829ba9..d3bc039 100644 --- a/Jenkinsfile +++ b/Jenkinsfile @@ -12,45 +12,22 @@ pipeline { environment { VENV = "${env.WORKSPACE}/generalresearch-venv" - SPECTRUM_CARER_VENV = "${env.WORKSPACE}/thl-spectrum-carer-venv" - GRLIQ_CARER_VENV = "${env.WORKSPACE}/grliq-carer-venv" - GR_CARER_VENV = "${env.WORKSPACE}/gr-carer-venv" - - INCITE_MOUNT_DIR = '/mnt/thl-incite' - TMP_DIR = "${env.WORKSPACE}/tmp" } stages { stage('python versions') { + matrix { axes { axis { name 'PYTHON_VERSION' - values 'python3.14' 'python3.13', 'python3.12', 'python3.11', 'python3.10' + values 'python3.14' 'python3.13', 'python3.12', 'python3.11' } } stages { - stage('Setup DB') { - script { - env.REDIS_DB = new Random().nextInt(1024).toString() - env.REDIS = "${env.REDIS}:6379/${env.REDIS_DB}" - env.THL_REDIS = "${env.THL_REDIS}:6379/${env.REDIS_DB}" - echo "Using THL Redis: ${env.REDIS}" - if (sh(script: "redis-cli -u ${env.REDIS} SET jenkins_lock 1 NX EX 3600", returnStdout: true).trim() != 'OK') - error('Redis already locked... aborting.') - } - 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.') - } - } - } - stage('Setup Git') { + stage('Setup') { steps { cleanWs() @@ -67,188 +44,32 @@ pipeline { url: 'ssh://code.g-r-l.com:6611/generalresearch'] ], ) - } - - dir("thl-spectrum:$PYTHON_VERSION/") { - checkout scmGit( - branches: [[name: env.BRANCH_NAME]], - extensions: [ cloneOption(shallow: true) ], - userRemoteConfigs: [ - [credentialsId: 'abdeb570-b708-44f3-b857-8a6b06ed9822', - url: 'ssh://code.g-r-l.com:6611/thl-marketplaces/thl-spectrum'] - ], - ) - } - dir("grliq:$PYTHON_VERSION/") { - checkout scmGit( - branches: [[name: env.BRANCH_NAME]], - extensions: [ cloneOption(shallow: true) ], - userRemoteConfigs: [ - [credentialsId: 'abdeb570-b708-44f3-b857-8a6b06ed9822', - url: 'ssh://code.g-r-l.com:6611/grl-iq'] - ], - ) - } - - dir("gr:$PYTHON_VERSION/") { - checkout scmGit( - branches: [[name: env.BRANCH_NAME]], - extensions: [ cloneOption(shallow: true) ], - userRemoteConfigs: [ - [credentialsId: 'abdeb570-b708-44f3-b857-8a6b06ed9822', - url: 'ssh://code.g-r-l.com:6611/general-research/gr-carer'] - ], - ) - } - } - } - - stage('Env & Migration') { - steps { - dir("generalresearch:$PYTHON_VERSION/") { sh "/usr/local/bin/$PYTHON_VERSION -m venv $VENV-$PYTHON_VERSION" sh "$VENV-$PYTHON_VERSION/bin/pip install -U setuptools wheel pip" sh "$VENV-$PYTHON_VERSION/bin/pip install -r requirements.txt" sh "$VENV-$PYTHON_VERSION/bin/pip install '.[django]'" - sh """ - export DB_NAME=${DB_NAME} - export DB_USER=${env.DB_USER} - export DB_PASSWORD=${env.DB_PASSWORD} - export DB_HOST=${env.DB_POSTGRESQL_HOST} - $VENV-$PYTHON_VERSION/bin/$PYTHON_VERSION -m generalresearch.thl_django.app.manage migrate - """ - } - - dir("thl-spectrum:$PYTHON_VERSION/") { - dir('carer') { - sh "/usr/local/bin/$PYTHON_VERSION -m venv $SPECTRUM_CARER_VENV-$PYTHON_VERSION" - sh "$SPECTRUM_CARER_VENV-$PYTHON_VERSION/bin/pip install -U setuptools wheel pip" - sh "$SPECTRUM_CARER_VENV-$PYTHON_VERSION/bin/pip install -r requirements.txt" - - sh """ - export DB_NAME=${SPECTRUM_DB_NAME} - $SPECTRUM_CARER_VENV-$PYTHON_VERSION/bin/$PYTHON_VERSION manage.py migrate --settings=carer.settings.unittest - """ - } - } - - dir("grliq:$PYTHON_VERSION/") { - dir('carer') { - sh "/usr/local/bin/$PYTHON_VERSION -m venv $GRLIQ_CARER_VENV-$PYTHON_VERSION" - sh "$GRLIQ_CARER_VENV-$PYTHON_VERSION/bin/pip install -U setuptools wheel pip" - sh "$GRLIQ_CARER_VENV-$PYTHON_VERSION/bin/pip install -r requirements.txt" - - sh """ - export DB_NAME=${GRLIQ_DB_NAME} - $GRLIQ_CARER_VENV-$PYTHON_VERSION/bin/$PYTHON_VERSION manage.py migrate --settings=carer.settings.unittest - """ - } - } - - dir("gr:$PYTHON_VERSION/") { - sh "/usr/local/bin/$PYTHON_VERSION -m venv $GR_CARER_VENV-$PYTHON_VERSION" - sh "$GR_CARER_VENV-$PYTHON_VERSION/bin/pip install -U setuptools wheel pip" - sh "$GR_CARER_VENV-$PYTHON_VERSION/bin/pip install -r requirements.txt" - - sh """ - export DB_NAME=${GR_DB_NAME} - $GR_CARER_VENV-$PYTHON_VERSION/bin/$PYTHON_VERSION manage.py migrate --settings=gr.settings.unittest - """ } } } stage('base') { - when { - expression { return true } - } - steps { - dir("generalresearch:$PYTHON_VERSION") { - sh "$VENV-$PYTHON_VERSION/bin/pytest -v tests/sql_helper.py" - } - } - } - - stage('models') { - when { - expression { return true } - } steps { dir("generalresearch:$PYTHON_VERSION") { - sh "$VENV-$PYTHON_VERSION/bin/pytest -v tests/models" + sh "$VENV-$PYTHON_VERSION/bin/pytest tests/models/gr/test_base.py -vs" } } } - stage('managers') { - steps { - dir("generalresearch:$PYTHON_VERSION") { - sh "$VENV-$PYTHON_VERSION/bin/pytest -v tests/managers" - } - } - } - - stage('wall_status_codes') { - steps { - dir("generalresearch:$PYTHON_VERSION") { - sh "$VENV-$PYTHON_VERSION/bin/pytest -v tests/wall_status_codes" - } - } - } - - stage('wxet') { - steps { - dir("generalresearch:$PYTHON_VERSION") { - sh "$VENV-$PYTHON_VERSION/bin/pytest -v tests/wxet" - } - } - } - - stage('grliq') { - steps { - dir("generalresearch:$PYTHON_VERSION") { - sh "$VENV-$PYTHON_VERSION/bin/pytest -v tests/grliq" - } - } - } - - stage('incite') { - steps { - dir("generalresearch:$PYTHON_VERSION") { - sh "$VENV-$PYTHON_VERSION/bin/pytest -v tests/incite" - } - } - } } } } } + post { always { echo 'One way or another, I have finished' deleteDir() /* clean up our workspace */ - sh """ - mariadb -h ${env.DB_MARIA_HOST} -u ${env.DB_USER} -p${env.DB_PASSWORD} --ssl=0 -e 'DROP DATABASE `${env.SPECTRUM_DB_NAME}`;' - """ - sh """ - PGPASSWORD=${env.DB_PASSWORD} psql -h ${env.DB_POSTGRESQL_HOST} -U ${env.DB_USER} -d postgres < datetime | None: + def check_start(cls, start: datetime | None) -> datetime | None: if start and start.microsecond != 0: raise ValueError("Collection.start must not have microseconds") return start @field_validator("offset") - def check_offset(cls, v: str | None, info: ValidationInfo): + def check_offset(cls, v: str | None): # pd.offsets.__all__ if v is None: # In MergeCollections, offset can be None diff --git a/generalresearch/managers/thl/ipinfo.py b/generalresearch/managers/thl/ipinfo.py index 98a9a32..496fe3e 100644 --- a/generalresearch/managers/thl/ipinfo.py +++ b/generalresearch/managers/thl/ipinfo.py @@ -6,6 +6,7 @@ from decimal import Decimal import faker import pymysql +from grip_client.enums import AccessType from more_itertools import chunked from psycopg import Cursor from pydantic import PositiveInt @@ -24,7 +25,6 @@ from generalresearch.models.thl.ipinfo import ( IPInformation, normalize_ip, ) -from generalresearch.models.thl.maxmind.definitions import UserType from generalresearch.pg_helper import PostgresConfig fake = faker.Faker() @@ -244,7 +244,7 @@ class IPInformationManager(PostgresManager): network: str | None = None, organization: str | None = None, static_ip_score: float | None = None, - user_type: UserType | None = None, + user_type: AccessType | None = None, postal_code: str | None = None, latitude: Decimal | None = None, longitude: Decimal | None = None, diff --git a/generalresearch/models/gr/__init__.py b/generalresearch/models/gr/__init__.py index 7e1516b..79f05d3 100644 --- a/generalresearch/models/gr/__init__.py +++ b/generalresearch/models/gr/__init__.py @@ -1,13 +1,13 @@ -from generalresearch.models.gr.authentication import GRToken, GRUser -from generalresearch.models.gr.business import Business -from generalresearch.models.gr.team import Team -from generalresearch.models.thl.finance import BusinessBalances -from generalresearch.models.thl.payout import BrokerageProductPayoutEvent -from generalresearch.models.thl.product import Product +# from generalresearch.models.gr.authentication import GRToken, GRUser +# from generalresearch.models.gr.business import Business +# from generalresearch.models.gr.team import Team +# from generalresearch.models.thl.finance import BusinessBalances +# from generalresearch.models.thl.payout import BrokerageProductPayoutEvent +# from generalresearch.models.thl.product import Product -_ = Business, Product, BrokerageProductPayoutEvent, BusinessBalances +# _ = Business, Product, BrokerageProductPayoutEvent, BusinessBalances -GRUser.model_rebuild() -GRToken.model_rebuild() -Business.model_rebuild() -Team.model_rebuild() +# GRUser.model_rebuild() +# GRToken.model_rebuild() +# Business.model_rebuild() +# Team.model_rebuild() diff --git a/generalresearch/models/gr/business.py b/generalresearch/models/gr/business.py index 11a5770..f6689a8 100644 --- a/generalresearch/models/gr/business.py +++ b/generalresearch/models/gr/business.py @@ -207,13 +207,13 @@ class Business(BaseModel): # Initialization is deferred until unless it's called # (see .prebuild_***()) - balance: BusinessBalances | None = Field(default=None, name="Business Balance") + balance: BusinessBalances | None = Field(default=None, title="Business Balance") payouts_total_str: str | None = Field(default=None) payouts_total: USDCent | None = Field(default=None) payouts: list[BusinessPayoutEvent] | None = Field( default=None, - name="Business Payouts", + title="Business Payouts", description=( "These are the ACH or Wire payments that were sent to the" "Business as a single amount, summed for all the Business" diff --git a/generalresearch/models/thl/session.py b/generalresearch/models/thl/session.py index 7121ea5..5653526 100644 --- a/generalresearch/models/thl/session.py +++ b/generalresearch/models/thl/session.py @@ -27,7 +27,6 @@ from generalresearch.models.custom_types import ( ) from generalresearch.models.legacy.bucket import Bucket from generalresearch.models.thl import ( - Product, decimal_to_int_cents, int_cents_to_decimal, ) @@ -47,6 +46,7 @@ if TYPE_CHECKING: 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("Wall") diff --git a/test_utils/conftest.py b/test_utils/conftest.py index 6894450..757f141 100644 --- a/test_utils/conftest.py +++ b/test_utils/conftest.py @@ -23,6 +23,18 @@ 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.') +# } + @pytest.fixture(scope="session") def env_file_path(pytestconfig: Config) -> Path: @@ -178,18 +190,18 @@ def gr_repo( repo_url = "ssh://code.g-r-l.com/general-research/gr-carer.git" _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: + if _ran.get(repo_url, False): + print(f"Already ran django_db_factory.{repo_url}") + return repo_path + + _ran[repo_url] = True + ssh_cmd = ( f"ssh -i {git_key_path} " "-o IdentitiesOnly=yes " @@ -206,11 +218,6 @@ def gr_repo( 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 @@ -236,7 +243,8 @@ def django_db_factory( if _ran.get(django_project, False): print(f"Already ran django_db_factory.{django_project}") - return + return postgres_instance + _ran[django_project] = True if "gr" in django_project: @@ -246,6 +254,8 @@ def django_db_factory( # 1. Bootstrapping Django settings if not django_settings.configured: + print(postgres_instance_dict) + django_settings.configure( DATABASES={ "default": { @@ -265,11 +275,13 @@ 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) + if "gr" in django_project: + call_command("makemigrations", "common", interactive=False) + call_command("migrate") # 3. Return the Dsn so the factory gives a way to connect diff --git a/test_utils/managers/cashout_methods.py b/test_utils/managers/cashout_methods.py index b201e8c..238cdda 100644 --- a/test_utils/managers/cashout_methods.py +++ b/test_utils/managers/cashout_methods.py @@ -1,6 +1,11 @@ +from __future__ import annotations + import random +from collections.abc import Callable from uuid import uuid4 +import pytest + from generalresearch.models.thl.wallet import Currency, PayoutType from generalresearch.models.thl.wallet.cashout_method import ( CashoutMethod, @@ -8,45 +13,55 @@ from generalresearch.models.thl.wallet.cashout_method import ( ) -def random_ext_id(base: str = "U02"): - suffix = random.randint(0, 99999) - return f"{base}{suffix:05d}" +@pytest.fixture(scope="session") +def random_ext_id_factory(base: str = "U02") -> Callable[..., str]: + + def _inner() -> str: + suffix = random.randint(0, 99999) + return f"{base}{suffix:05d}" + return _inner -EXAMPLE_TANGO_CASHOUT_METHODS = [ - CashoutMethod( - id=uuid4().hex, - last_updated="2021-06-23T20:45:38.239182Z", - is_live=True, - type=PayoutType.TANGO, - ext_id=random_ext_id(), - name="Safeway eGift Card $25", - data=TangoCashoutMethodData( - value_type="fixed", countries=["US"], utid=random_ext_id() + +@pytest.fixture(scope="session") +def example_tango_cashout_methods( + random_ext_id_factory: Callable[..., str], +) -> list[CashoutMethod]: + return [ + CashoutMethod( + id=uuid4().hex, + last_updated="2021-06-23T20:45:38.239182Z", + is_live=True, + type=PayoutType.TANGO, + ext_id=random_ext_id_factory(), + name="Safeway eGift Card $25", + data=TangoCashoutMethodData( + value_type="fixed", countries=["US"], utid=random_ext_id_factory() + ), + user=None, + image_url="https://d30s7yzk2az89n.cloudfront.net/images/brands/b694446-1200w-326ppi.png", + original_currency=Currency.USD, + min_value=2500, + max_value=2500, ), - user=None, - image_url="https://d30s7yzk2az89n.cloudfront.net/images/brands/b694446-1200w-326ppi.png", - original_currency=Currency.USD, - min_value=2500, - max_value=2500, - ), - CashoutMethod( - id=uuid4().hex, - last_updated="2021-06-23T20:45:38.239182Z", - is_live=True, - type=PayoutType.TANGO, - ext_id=random_ext_id(), - name="Amazon.it Gift Certificate", - data=TangoCashoutMethodData( - value_type="variable", countries=["IT"], utid="U006961" + CashoutMethod( + id=uuid4().hex, + last_updated="2021-06-23T20:45:38.239182Z", + is_live=True, + type=PayoutType.TANGO, + ext_id=random_ext_id_factory(), + name="Amazon.it Gift Certificate", + data=TangoCashoutMethodData( + value_type="variable", countries=["IT"], utid="U006961" + ), + user=None, + image_url="https://d30s7yzk2az89n.cloudfront.net/images/brands/b405753-1200w-326ppi.png", + original_currency=Currency.EUR, + min_value=1, + max_value=10000, ), - user=None, - image_url="https://d30s7yzk2az89n.cloudfront.net/images/brands/b405753-1200w-326ppi.png", - original_currency=Currency.EUR, - min_value=1, - max_value=10000, - ), -] + ] + # AMT_ASSIGNMENT_CASHOUT_METHOD = CashoutMethod( # id=uuid4().hex, diff --git a/test_utils/managers/conftest.py b/test_utils/managers/conftest.py index b03a646..9c6a1a7 100644 --- a/test_utils/managers/conftest.py +++ b/test_utils/managers/conftest.py @@ -32,12 +32,10 @@ from generalresearch.managers.thl.userhealth import ( 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 test_utils.managers.cashout_methods import ( - EXAMPLE_TANGO_CASHOUT_METHODS, -) # === THL === @@ -172,12 +170,13 @@ def delete_cashoutmethod_db(thl_web_rw: PostgresConfig) -> Callable[..., None]: def setup_cashoutmethod_db( cashout_method_manager: CashoutMethodManager, delete_cashoutmethod_db: Callable[..., None], + example_tango_cashout_methods: list[CashoutMethod], ) -> Callable[..., None]: def _inner(): delete_cashoutmethod_db() - for x in EXAMPLE_TANGO_CASHOUT_METHODS: + for x in example_tango_cashout_methods: cashout_method_manager.create(x) # TODO: convert these ids into instances to use. diff --git a/test_utils/managers/gr/conftest.py b/test_utils/managers/gr/conftest.py index 4da8fe3..69f3e9a 100644 --- a/test_utils/managers/gr/conftest.py +++ b/test_utils/managers/gr/conftest.py @@ -56,9 +56,11 @@ def gr_redis_config(settings: GRLBaseSettings) -> RedisConfig: @pytest.fixture(scope="session") def gr_db(django_db_factory: Callable[..., PostgresDsn]) -> PostgresConfig: + _dsn = django_db_factory("gr.common") + print("DDDD:", _dsn) return PostgresConfig( - dsn=django_db_factory("gr_carer"), + dsn=_dsn, connect_timeout=1, statement_timeout=5, ) diff --git a/test_utils/models/thl/conftest.py b/test_utils/models/thl/conftest.py index 907306d..fc57c73 100644 --- a/test_utils/models/thl/conftest.py +++ b/test_utils/models/thl/conftest.py @@ -4,13 +4,13 @@ from collections.abc import Callable from datetime import UTC, datetime 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 from uuid import uuid4 import faker import pytest +from grip_client.enums import AccessType from pydantic import PositiveInt from generalresearch.managers.thl.ipinfo import IPGeonameManager, IPInformationManager @@ -30,7 +30,7 @@ from generalresearch.models.legacy.bucket import Bucket from generalresearch.models.thl.definitions import ( PayoutStatus, ) -from generalresearch.models.thl.ipinfo import IPGeoname, IPInformation, UserType +from generalresearch.models.thl.ipinfo import IPGeoname, IPInformation from generalresearch.models.thl.payout import UserPayoutEvent from generalresearch.models.thl.product import ( PayoutConfig, @@ -118,7 +118,7 @@ def wall_factory( session = session_factory() session_id = session.id - source = source or rchoice(list(Source)) + 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 @@ -284,7 +284,7 @@ def ipinformation_factory( network: str | None = None, organization: str | None = None, static_ip_score: float | None = None, - user_type: UserType | None = None, + user_type: AccessType | None = None, postal_code: str | None = None, latitude: Decimal | None = None, longitude: Decimal | None = None, @@ -426,8 +426,8 @@ def auditlog_factory(audit_log_manager: AuditLogManager): return audit_log_manager.create( user_id=user_id, - level=level or rchoice(list(AuditLogLevel)), - event_type=event_type or rchoice(list(event_types)), + level=level or rand_choice(list(AuditLogLevel)), + event_type=event_type or rand_choice(list(event_types)), event_msg=event_msg, event_value=event_value, ) diff --git a/tests/managers/thl/test_cashout_method.py b/tests/managers/thl/test_cashout_method.py index 451d3e0..ca85c6b 100644 --- a/tests/managers/thl/test_cashout_method.py +++ b/tests/managers/thl/test_cashout_method.py @@ -12,12 +12,10 @@ 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 test_utils.managers.cashout_methods import ( - EXAMPLE_TANGO_CASHOUT_METHODS, -) class TestTangoCashoutMethods: @@ -26,13 +24,14 @@ class TestTangoCashoutMethods: self, cashout_method_manager: CashoutMethodManager, setup_cashoutmethod_db: Callable[..., None], + example_tango_cashout_methods: list[CashoutMethod], ): setup_cashoutmethod_db() res = cashout_method_manager.filter(payout_types=[PayoutType.TANGO]) assert len(res) == 2 cm = next(x for x in res if x.ext_id == "U025035") - assert EXAMPLE_TANGO_CASHOUT_METHODS[0] == cm + assert example_tango_cashout_methods[0] == cm def test_user( self, diff --git a/tests/models/gr/test_base.py b/tests/models/gr/test_base.py index 412fa52..a066fa1 100644 --- a/tests/models/gr/test_base.py +++ b/tests/models/gr/test_base.py @@ -15,8 +15,6 @@ class TestGRPostgresDjangoCreation: 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( @@ -37,16 +35,15 @@ class TestGRPostgresDjangoCreation: django_db_factory: Callable[..., None], ): - dsn = django_db_factory("gr") + dsn = django_db_factory("gr.common") assert isinstance(dsn, PostgresDsn) def test_django_tables(self, gr_db: PostgresConfig): res = gr_db.execute_sql_query(query=""" - SELECT COUNT(*) + SELECT COUNT(*) FROM information_schema.tables WHERE table_schema = 'public'; """) print(res) assert len(res) == 1 - assert res[0]["count"] == 56 - assert True + assert res[0]["count"] == 10 -- cgit v1.2.3 From d19438ec4ccbbe4415c286c9ae89e3e5706ac553 Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Mon, 31 Aug 2026 19:20:32 -0700 Subject: TYPE_CHECKING on models + managers --- generalresearch/config.py | 2 + generalresearch/grliq/managers/forensic_data.py | 10 +- generalresearch/grliq/models/forensic_data.py | 8 +- generalresearch/managers/base.py | 8 +- generalresearch/managers/cint/profiling.py | 5 +- generalresearch/managers/criteria.py | 5 +- generalresearch/managers/dynata/profiling.py | 5 +- generalresearch/managers/events.py | 10 +- generalresearch/managers/gr/authentication.py | 10 +- generalresearch/managers/gr/business.py | 24 +- generalresearch/managers/gr/team.py | 11 +- generalresearch/managers/innovate/profiling.py | 5 +- generalresearch/managers/leaderboard/tasks.py | 5 +- generalresearch/managers/lucid/profiling.py | 8 +- generalresearch/managers/marketplace/user_pid.py | 5 +- generalresearch/managers/morning/profiling.py | 5 +- generalresearch/managers/network/label.py | 11 +- generalresearch/managers/network/mtr.py | 6 +- generalresearch/managers/network/nmap.py | 6 +- generalresearch/managers/network/rdns.py | 6 +- generalresearch/managers/network/tool_run.py | 10 +- generalresearch/managers/pollfish/profiling.py | 5 +- generalresearch/managers/precision/profiling.py | 5 +- generalresearch/managers/prodege/profiling.py | 5 +- generalresearch/managers/repdata/profiling.py | 5 +- generalresearch/managers/repdata/survey.py | 5 +- generalresearch/managers/sago/profiling.py | 5 +- generalresearch/managers/spectrum/profiling.py | 5 +- generalresearch/managers/survey.py | 5 +- generalresearch/managers/thl/buyer.py | 7 +- generalresearch/managers/thl/cashout_method.py | 12 +- generalresearch/managers/thl/category.py | 10 +- generalresearch/managers/thl/contest_manager.py | 28 +- generalresearch/managers/thl/ipinfo.py | 13 +- .../managers/thl/ledger_manager/conditions.py | 17 +- .../managers/thl/ledger_manager/ledger.py | 13 +- .../managers/thl/ledger_manager/thl_ledger.py | 33 +- generalresearch/managers/thl/payout.py | 20 +- generalresearch/managers/thl/product.py | 13 +- generalresearch/managers/thl/profiling/uqa.py | 5 +- generalresearch/managers/thl/profiling/user_upk.py | 12 +- generalresearch/managers/thl/session.py | 29 +- generalresearch/managers/thl/survey.py | 8 +- generalresearch/managers/thl/survey_penalty.py | 23 +- generalresearch/managers/thl/task_adjustment.py | 11 +- generalresearch/managers/thl/user_compensate.py | 12 +- .../managers/thl/user_manager/__init__.py | 5 +- .../thl/user_manager/mysql_user_manager.py | 7 +- .../managers/thl/user_manager/rate_limit.py | 5 +- .../managers/thl/user_manager/user_manager.py | 7 +- generalresearch/managers/thl/userhealth.py | 13 +- generalresearch/managers/thl/wall.py | 21 +- generalresearch/managers/thl/wallet/__init__.py | 36 ++- generalresearch/managers/thl/wallet/approve.py | 18 +- generalresearch/managers/thl/wallet/tango.py | 21 +- generalresearch/models/admin/request.py | 5 +- generalresearch/models/cint/question.py | 4 +- generalresearch/models/cint/survey.py | 16 +- generalresearch/models/cint/task_collection.py | 6 +- generalresearch/models/dynata/question.py | 6 +- generalresearch/models/dynata/survey.py | 21 +- generalresearch/models/dynata/task_collection.py | 6 +- generalresearch/models/events.py | 27 +- generalresearch/models/gr/authentication.py | 6 +- generalresearch/models/gr/business.py | 53 ++-- generalresearch/models/gr/team.py | 67 ++-- generalresearch/models/innovate/question.py | 2 +- generalresearch/models/innovate/survey.py | 21 +- generalresearch/models/innovate/task_collection.py | 6 +- generalresearch/models/legacy/bucket.py | 14 +- generalresearch/models/legacy/offerwall.py | 38 ++- generalresearch/models/legacy/questions.py | 12 +- generalresearch/models/lucid/question.py | 2 +- generalresearch/models/lucid/survey.py | 18 +- generalresearch/models/morning/question.py | 6 +- generalresearch/models/morning/survey.py | 26 +- generalresearch/models/morning/task_collection.py | 6 +- generalresearch/models/network/label.py | 9 +- generalresearch/models/network/mtr/command.py | 2 +- generalresearch/models/network/mtr/execute.py | 12 +- generalresearch/models/network/mtr/parser.py | 5 +- generalresearch/models/network/mtr/result.py | 9 +- generalresearch/models/network/nmap/command.py | 2 +- generalresearch/models/network/nmap/execute.py | 12 +- generalresearch/models/network/nmap/result.py | 6 +- generalresearch/models/network/rdns/command.py | 2 +- generalresearch/models/network/rdns/execute.py | 5 +- generalresearch/models/network/rdns/parser.py | 5 +- generalresearch/models/network/rdns/result.py | 4 +- generalresearch/models/network/tool_run.py | 31 +- generalresearch/models/network/tool_run_command.py | 6 +- generalresearch/models/precision/question.py | 2 +- generalresearch/models/precision/survey.py | 21 +- .../models/precision/task_collection.py | 6 +- generalresearch/models/prodege/question.py | 4 +- generalresearch/models/prodege/survey.py | 24 +- generalresearch/models/prodege/task_collection.py | 6 +- generalresearch/models/repdata/question.py | 2 +- generalresearch/models/repdata/survey.py | 14 +- generalresearch/models/repdata/task_collection.py | 6 +- generalresearch/models/sago/question.py | 2 +- generalresearch/models/sago/survey.py | 20 +- generalresearch/models/sago/task_collection.py | 6 +- generalresearch/models/spectrum/question.py | 4 +- generalresearch/models/spectrum/survey.py | 18 +- generalresearch/models/spectrum/task_collection.py | 6 +- generalresearch/models/thl/category.py | 5 +- generalresearch/models/thl/contest/__init__.py | 10 +- generalresearch/models/thl/contest/contest.py | 12 +- .../models/thl/contest/contest_entry.py | 10 +- generalresearch/models/thl/contest/leaderboard.py | 10 +- generalresearch/models/thl/contest/milestone.py | 6 +- generalresearch/models/thl/contest/raffle.py | 12 +- generalresearch/models/thl/demographics.py | 3 +- generalresearch/models/thl/finance.py | 4 +- generalresearch/models/thl/ipinfo.py | 11 +- generalresearch/models/thl/leaderboard.py | 7 +- generalresearch/models/thl/ledger.py | 19 +- generalresearch/models/thl/offerwall/base.py | 26 +- generalresearch/models/thl/offerwall/cache.py | 19 +- generalresearch/models/thl/payout.py | 25 +- generalresearch/models/thl/product.py | 16 +- .../models/thl/profiling/marketplace.py | 21 +- generalresearch/models/thl/profiling/question.py | 17 +- .../models/thl/profiling/upk_property.py | 7 +- .../models/thl/profiling/upk_question.py | 6 +- .../models/thl/profiling/upk_question_answer.py | 14 +- generalresearch/models/thl/profiling/user_info.py | 15 +- .../models/thl/profiling/user_question_answer.py | 13 +- generalresearch/models/thl/report_task.py | 5 +- generalresearch/models/thl/session.py | 23 +- generalresearch/models/thl/soft_pair.py | 12 +- generalresearch/models/thl/survey/__init__.py | 21 +- generalresearch/models/thl/survey/buyer.py | 14 +- generalresearch/models/thl/survey/model.py | 25 +- generalresearch/models/thl/survey/penalty.py | 13 +- .../models/thl/survey/task_collection.py | 4 +- generalresearch/models/thl/task_adjustment.py | 8 +- generalresearch/models/thl/task_status.py | 33 +- generalresearch/models/thl/user.py | 12 +- generalresearch/models/thl/user_iphistory.py | 17 +- generalresearch/models/thl/user_profile.py | 10 +- generalresearch/models/thl/user_quality_event.py | 10 +- generalresearch/models/thl/user_streak.py | 5 +- .../models/thl/wallet/cashout_method.py | 23 +- generalresearch/models/thl/wallet/payout.py | 12 +- generalresearch/models/thl/wallet/user_wallet.py | 7 +- generalresearch/wall_status_codes/cint.py | 6 +- test_utils/conftest.py | 2 +- test_utils/managers/gr/conftest.py | 48 ++- test_utils/models/conftest.py | 37 --- test_utils/models/gr/conftest.py | 56 +++- tests/models/gr/test_business.py | 352 +++++++++++---------- 153 files changed, 1319 insertions(+), 947 deletions(-) (limited to 'test_utils/models') diff --git a/generalresearch/config.py b/generalresearch/config.py index 73af565..2777aa4 100644 --- a/generalresearch/config.py +++ b/generalresearch/config.py @@ -53,6 +53,8 @@ class GRLBaseSettings(BaseSettings): testing_postgres_user: str | None = Field(default=None) testing_postgres_pass: str | None = Field(default=None) + testing_redis: InternalHostname | None = Field(default=None) + git_creds: str | None = Field(default=None) # --- diff --git a/generalresearch/grliq/managers/forensic_data.py b/generalresearch/grliq/managers/forensic_data.py index c1eac37..0f8534c 100644 --- a/generalresearch/grliq/managers/forensic_data.py +++ b/generalresearch/grliq/managers/forensic_data.py @@ -2,7 +2,7 @@ from __future__ import annotations from collections.abc import Collection from datetime import datetime -from typing import Any +from typing import TYPE_CHECKING, Any from psycopg import sql from pydantic import NonNegativeInt, PositiveInt @@ -14,9 +14,11 @@ from generalresearch.grliq.models.forensic_result import ( GrlIqForensicCategoryResult, Phase, ) -from generalresearch.models.custom_types import UUIDStr -from generalresearch.models.thl.user import User -from generalresearch.pg_helper import PostgresConfig + +if TYPE_CHECKING: + from generalresearch.models.custom_types import UUIDStr + from generalresearch.models.thl.user import User + from generalresearch.pg_helper import PostgresConfig class GrlIqDataManager: diff --git a/generalresearch/grliq/models/forensic_data.py b/generalresearch/grliq/models/forensic_data.py index 69b1760..6a07774 100644 --- a/generalresearch/grliq/models/forensic_data.py +++ b/generalresearch/grliq/models/forensic_data.py @@ -6,7 +6,7 @@ from collections import Counter from datetime import UTC, datetime, timedelta from enum import StrEnum from functools import cached_property -from typing import Annotated, Any, Literal, Self +from typing import TYPE_CHECKING, Annotated, Any, Literal, Self from uuid import uuid4 import pycountry @@ -53,8 +53,10 @@ from generalresearch.models.custom_types import ( IPvAnyAddressStr, UUIDStr, ) -from generalresearch.models.thl.ipinfo import GeoIPInformation -from generalresearch.models.thl.session import Session + +if TYPE_CHECKING: + from generalresearch.models.thl.ipinfo import GeoIPInformation + from generalresearch.models.thl.session import Session fake = Faker() diff --git a/generalresearch/managers/base.py b/generalresearch/managers/base.py index ba42e9e..2227d21 100644 --- a/generalresearch/managers/base.py +++ b/generalresearch/managers/base.py @@ -3,10 +3,12 @@ from __future__ import annotations from collections.abc import Collection from contextlib import nullcontext from enum import Enum +from typing import TYPE_CHECKING -from generalresearch.pg_helper import PostgresConfig -from generalresearch.redis_helper import RedisConfig -from generalresearch.sql_helper import SqlHelper +if TYPE_CHECKING: + from generalresearch.pg_helper import PostgresConfig + from generalresearch.redis_helper import RedisConfig + from generalresearch.sql_helper import SqlHelper class Permission(int, Enum): diff --git a/generalresearch/managers/cint/profiling.py b/generalresearch/managers/cint/profiling.py index 9216aa5..c550632 100644 --- a/generalresearch/managers/cint/profiling.py +++ b/generalresearch/managers/cint/profiling.py @@ -2,9 +2,12 @@ from __future__ import annotations import json from collections.abc import Collection +from typing import TYPE_CHECKING from generalresearch.models.cint.question import CintQuestion -from generalresearch.sql_helper import SqlHelper + +if TYPE_CHECKING: + from generalresearch.sql_helper import SqlHelper def get_profiling_library( diff --git a/generalresearch/managers/criteria.py b/generalresearch/managers/criteria.py index b8ae6ac..760afd3 100644 --- a/generalresearch/managers/criteria.py +++ b/generalresearch/managers/criteria.py @@ -3,11 +3,14 @@ from __future__ import annotations from abc import ABC from collections.abc import Collection from datetime import UTC, datetime +from typing import TYPE_CHECKING from more_itertools import chunked from generalresearch.managers.base import SqlManager -from generalresearch.models.thl.survey import MarketplaceCondition + +if TYPE_CHECKING: + from generalresearch.models.thl.survey import MarketplaceCondition DB_FIELDS = [ "hash", diff --git a/generalresearch/managers/dynata/profiling.py b/generalresearch/managers/dynata/profiling.py index 661bdad..b76ecc7 100644 --- a/generalresearch/managers/dynata/profiling.py +++ b/generalresearch/managers/dynata/profiling.py @@ -2,9 +2,12 @@ from __future__ import annotations import json from collections.abc import Collection +from typing import TYPE_CHECKING from generalresearch.models.dynata.question import DynataQuestion -from generalresearch.sql_helper import SqlHelper + +if TYPE_CHECKING: + from generalresearch.sql_helper import SqlHelper def get_profiling_library( diff --git a/generalresearch/managers/events.py b/generalresearch/managers/events.py index 6779104..c43a020 100644 --- a/generalresearch/managers/events.py +++ b/generalresearch/managers/events.py @@ -13,14 +13,12 @@ from redis.client import PubSub, Redis from generalresearch.incite.base import LOG from generalresearch.managers.base import RedisManager from generalresearch.models import Source -from generalresearch.models.custom_types import UUIDStr from generalresearch.models.events import ( AggregateBySource, EventEnvelope, EventMessage, EventType, MaxGaugeBySource, - ServerToClientMessage, ServerToClientMessageAdapter, SessionEnterPayload, SessionFinishPayload, @@ -30,11 +28,15 @@ from generalresearch.models.events import ( TaskStatsSnapshot, ) from generalresearch.models.thl.definitions import Status -from generalresearch.models.thl.session import Session, Wall -from generalresearch.models.thl.user import User if TYPE_CHECKING: from influxdb import InfluxDBClient + + from generalresearch.models.custom_types import UUIDStr + from generalresearch.models.events import ServerToClientMessage + from generalresearch.models.thl.session import Session, Wall + from generalresearch.models.thl.user import User + else: InfluxDBClient = object diff --git a/generalresearch/managers/gr/authentication.py b/generalresearch/managers/gr/authentication.py index 851b88a..ca56467 100644 --- a/generalresearch/managers/gr/authentication.py +++ b/generalresearch/managers/gr/authentication.py @@ -10,14 +10,14 @@ from psycopg import sql from pydantic import AnyHttpUrl, PositiveInt from generalresearch.managers.base import PostgresManager, PostgresManagerWithRedis -from generalresearch.models.custom_types import UUIDStr -from generalresearch.pg_helper import PostgresConfig -from generalresearch.redis_helper import RedisConfig - -LOG = logging.getLogger("gr") if TYPE_CHECKING: + from generalresearch.models.custom_types import UUIDStr from generalresearch.models.gr.authentication import GRToken, GRUser + from generalresearch.pg_helper import PostgresConfig + from generalresearch.redis_helper import RedisConfig + +LOG = logging.getLogger("gr") class GRUserManager(PostgresManagerWithRedis): diff --git a/generalresearch/managers/gr/business.py b/generalresearch/managers/gr/business.py index 4da0e7f..ef26f30 100644 --- a/generalresearch/managers/gr/business.py +++ b/generalresearch/managers/gr/business.py @@ -11,14 +11,16 @@ from generalresearch.managers.base import ( PostgresManager, PostgresManagerWithRedis, ) -from generalresearch.models.custom_types import UUIDStr +from generalresearch.models.gr.business import ( + Business, + BusinessBankAccount, + BusinessType, +) if TYPE_CHECKING: + from generalresearch.models.custom_types import UUIDStr from generalresearch.models.gr.business import ( - Business, BusinessAddress, - BusinessBankAccount, - BusinessType, TransferMethod, ) from generalresearch.models.gr.team import Team @@ -36,8 +38,6 @@ class BusinessBankAccountManager(PostgresManager): iban: str | None = None, swift: str | None = None, ) -> BusinessBankAccount: - from generalresearch.models.gr.business import BusinessBankAccount - ba = BusinessBankAccount.model_validate( { "business_id": business_id, @@ -73,7 +73,6 @@ class BusinessBankAccountManager(PostgresManager): return ba def get_by_business_id(self, business_id: UUIDStr) -> list[BusinessBankAccount]: - from generalresearch.models.gr.business import BusinessBankAccount with self.pg_config.make_connection() as conn, conn.cursor() as c: c.execute( @@ -192,11 +191,7 @@ class BusinessManager(PostgresManagerWithRedis): """ Behavior: does this raise on duplicate? """ - from generalresearch.models.gr.business import ( - Business, - BusinessType, - ) - + # Business.model_rebuild() business = Business.model_validate( { "uuid": uuid or uuid4().hex, @@ -281,7 +276,6 @@ class BusinessManager(PostgresManagerWithRedis): res = c.fetchall() response = [] - from generalresearch.models.gr.business import Business for i in res: # i["contact"] = BusinessContact.model_validate(i) @@ -370,8 +364,6 @@ class BusinessManager(PostgresManagerWithRedis): self, business_uuid: UUIDStr, ) -> Business | None: - from generalresearch.models.gr.business import Business - assert UUID(hex=business_uuid).hex == business_uuid with self.pg_config.make_connection() as conn, conn.cursor() as c: @@ -397,8 +389,6 @@ class BusinessManager(PostgresManagerWithRedis): return Business.model_validate(data) def get_by_id(self, business_id: PositiveInt) -> Business | None: - from generalresearch.models.gr.business import Business - assert isinstance(business_id, int) with self.pg_config.make_connection() as conn, conn.cursor() as c: diff --git a/generalresearch/managers/gr/team.py b/generalresearch/managers/gr/team.py index d04370b..3283467 100644 --- a/generalresearch/managers/gr/team.py +++ b/generalresearch/managers/gr/team.py @@ -11,15 +11,16 @@ from generalresearch.managers.base import ( PostgresManager, PostgresManagerWithRedis, ) -from generalresearch.models.custom_types import UUIDStr -from generalresearch.models.gr.team import Membership, MembershipPrivilege +from generalresearch.models.gr.team import ( + Membership, + MembershipPrivilege, +) if TYPE_CHECKING: + from generalresearch.models.custom_types import UUIDStr from generalresearch.models.gr.authentication import GRUser from generalresearch.models.gr.business import Business - from generalresearch.models.gr.team import ( - Team, - ) + from generalresearch.models.gr.team import Team class MembershipManager(PostgresManager): diff --git a/generalresearch/managers/innovate/profiling.py b/generalresearch/managers/innovate/profiling.py index bfa2685..0f32999 100644 --- a/generalresearch/managers/innovate/profiling.py +++ b/generalresearch/managers/innovate/profiling.py @@ -2,9 +2,12 @@ from __future__ import annotations import json from collections.abc import Collection +from typing import TYPE_CHECKING from generalresearch.models.innovate.question import InnovateQuestion -from generalresearch.sql_helper import SqlHelper + +if TYPE_CHECKING: + from generalresearch.sql_helper import SqlHelper def get_profiling_library( diff --git a/generalresearch/managers/leaderboard/tasks.py b/generalresearch/managers/leaderboard/tasks.py index 072a1aa..9e8dc9f 100644 --- a/generalresearch/managers/leaderboard/tasks.py +++ b/generalresearch/managers/leaderboard/tasks.py @@ -1,4 +1,5 @@ import logging +from typing import TYPE_CHECKING from redis import Redis @@ -7,7 +8,9 @@ from generalresearch.models.thl.leaderboard import ( LeaderboardCode, LeaderboardFrequency, ) -from generalresearch.models.thl.session import Session + +if TYPE_CHECKING: + from generalresearch.models.thl.session import Session logger = logging.getLogger() diff --git a/generalresearch/managers/lucid/profiling.py b/generalresearch/managers/lucid/profiling.py index fdd2d52..5937a59 100644 --- a/generalresearch/managers/lucid/profiling.py +++ b/generalresearch/managers/lucid/profiling.py @@ -2,12 +2,16 @@ from __future__ import annotations import json from collections.abc import Collection +from typing import TYPE_CHECKING from pydantic import ValidationError -from generalresearch.decorators import LOG from generalresearch.models.lucid.question import LucidQuestion, LucidQuestionType -from generalresearch.sql_helper import SqlHelper + +if TYPE_CHECKING: + from generalresearch.sql_helper import SqlHelper + +from generalresearch.decorators import LOG def get_profiling_library( diff --git a/generalresearch/managers/marketplace/user_pid.py b/generalresearch/managers/marketplace/user_pid.py index 15d8a19..fe24d38 100644 --- a/generalresearch/managers/marketplace/user_pid.py +++ b/generalresearch/managers/marketplace/user_pid.py @@ -2,11 +2,14 @@ from __future__ import annotations from abc import ABC from collections.abc import Collection +from typing import TYPE_CHECKING from uuid import UUID from generalresearch.managers.base import SqlManager from generalresearch.models import Source -from generalresearch.sql_helper import SqlHelper + +if TYPE_CHECKING: + from generalresearch.sql_helper import SqlHelper class UserPidManager(SqlManager, ABC): diff --git a/generalresearch/managers/morning/profiling.py b/generalresearch/managers/morning/profiling.py index 01f99f3..7335e6b 100644 --- a/generalresearch/managers/morning/profiling.py +++ b/generalresearch/managers/morning/profiling.py @@ -2,9 +2,12 @@ from __future__ import annotations import json from collections.abc import Collection +from typing import TYPE_CHECKING from generalresearch.models.morning.question import MorningQuestion -from generalresearch.sql_helper import SqlHelper + +if TYPE_CHECKING: + from generalresearch.sql_helper import SqlHelper def get_profiling_library( diff --git a/generalresearch/managers/network/label.py b/generalresearch/managers/network/label.py index f0ba9f7..cec59ad 100644 --- a/generalresearch/managers/network/label.py +++ b/generalresearch/managers/network/label.py @@ -2,6 +2,7 @@ from __future__ import annotations from collections.abc import Collection from datetime import UTC, datetime, timedelta +from typing import TYPE_CHECKING from psycopg import sql from pydantic import IPvAnyNetwork, TypeAdapter @@ -9,10 +10,16 @@ from pydantic import IPvAnyNetwork, TypeAdapter from generalresearch.managers.base import PostgresManager from generalresearch.models.custom_types import ( AwareDatetimeISO, - IPvAnyAddressStr, IPvAnyNetworkStr, ) -from generalresearch.models.network.label import IPLabel, IPLabelKind, IPLabelSource +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 class IPLabelManager(PostgresManager): diff --git a/generalresearch/managers/network/mtr.py b/generalresearch/managers/network/mtr.py index 179b8a9..7b79d96 100644 --- a/generalresearch/managers/network/mtr.py +++ b/generalresearch/managers/network/mtr.py @@ -1,9 +1,13 @@ from __future__ import annotations +from typing import TYPE_CHECKING + from psycopg import Cursor, sql from generalresearch.managers.base import PostgresManager -from generalresearch.models.network.tool_run import MTRRun + +if TYPE_CHECKING: + from generalresearch.models.network.tool_run import MTRRun class MTRRunManager(PostgresManager): diff --git a/generalresearch/managers/network/nmap.py b/generalresearch/managers/network/nmap.py index 574bce1..96c6009 100644 --- a/generalresearch/managers/network/nmap.py +++ b/generalresearch/managers/network/nmap.py @@ -1,9 +1,13 @@ from __future__ import annotations +from typing import TYPE_CHECKING + from psycopg import Cursor, sql from generalresearch.managers.base import PostgresManager -from generalresearch.models.network.tool_run import NmapRun + +if TYPE_CHECKING: + from generalresearch.models.network.tool_run import NmapRun class NmapRunManager(PostgresManager): diff --git a/generalresearch/managers/network/rdns.py b/generalresearch/managers/network/rdns.py index 95a1381..1800364 100644 --- a/generalresearch/managers/network/rdns.py +++ b/generalresearch/managers/network/rdns.py @@ -1,9 +1,13 @@ from __future__ import annotations +from typing import TYPE_CHECKING + from psycopg import Cursor from generalresearch.managers.base import PostgresManager -from generalresearch.models.network.tool_run import RDNSRun + +if TYPE_CHECKING: + from generalresearch.models.network.tool_run import RDNSRun class RDNSRunManager(PostgresManager): diff --git a/generalresearch/managers/network/tool_run.py b/generalresearch/managers/network/tool_run.py index 73afc13..ec06305 100644 --- a/generalresearch/managers/network/tool_run.py +++ b/generalresearch/managers/network/tool_run.py @@ -1,10 +1,11 @@ from __future__ import annotations from collections.abc import Collection +from typing import TYPE_CHECKING from psycopg import Cursor, sql -from generalresearch.managers.base import Permission, PostgresManager +from generalresearch.managers.base import PostgresManager from generalresearch.managers.network.mtr import MTRRunManager from generalresearch.managers.network.nmap import NmapRunManager from generalresearch.managers.network.rdns import RDNSRunManager @@ -13,10 +14,13 @@ from generalresearch.models.network.tool_run import ( MTRRun, NmapRun, RDNSRun, - ToolName, ToolRun, ) -from generalresearch.pg_helper import PostgresConfig + +if TYPE_CHECKING: + from generalresearch.managers.base import Permission + from generalresearch.models.network.tool_run import ToolName + from generalresearch.pg_helper import PostgresConfig class ToolRunManager(PostgresManager): diff --git a/generalresearch/managers/pollfish/profiling.py b/generalresearch/managers/pollfish/profiling.py index daf529b..43afd30 100644 --- a/generalresearch/managers/pollfish/profiling.py +++ b/generalresearch/managers/pollfish/profiling.py @@ -2,9 +2,12 @@ from __future__ import annotations import json from collections.abc import Collection +from typing import TYPE_CHECKING from generalresearch.models.pollfish.question import PollfishQuestion -from generalresearch.sql_helper import SqlHelper + +if TYPE_CHECKING: + from generalresearch.sql_helper import SqlHelper def get_profiling_library( diff --git a/generalresearch/managers/precision/profiling.py b/generalresearch/managers/precision/profiling.py index 449fd25..813542c 100644 --- a/generalresearch/managers/precision/profiling.py +++ b/generalresearch/managers/precision/profiling.py @@ -2,9 +2,12 @@ from __future__ import annotations import json from collections.abc import Collection +from typing import TYPE_CHECKING from generalresearch.models.precision.question import PrecisionQuestion -from generalresearch.sql_helper import SqlHelper + +if TYPE_CHECKING: + from generalresearch.sql_helper import SqlHelper def get_profiling_library( diff --git a/generalresearch/managers/prodege/profiling.py b/generalresearch/managers/prodege/profiling.py index 54a7b57..cf77fca 100644 --- a/generalresearch/managers/prodege/profiling.py +++ b/generalresearch/managers/prodege/profiling.py @@ -2,9 +2,12 @@ from __future__ import annotations import json from collections.abc import Collection +from typing import TYPE_CHECKING from generalresearch.models.prodege.question import ProdegeQuestion -from generalresearch.sql_helper import SqlHelper + +if TYPE_CHECKING: + from generalresearch.sql_helper import SqlHelper def get_profiling_library( diff --git a/generalresearch/managers/repdata/profiling.py b/generalresearch/managers/repdata/profiling.py index 4b97abd..ec30d78 100644 --- a/generalresearch/managers/repdata/profiling.py +++ b/generalresearch/managers/repdata/profiling.py @@ -2,9 +2,12 @@ from __future__ import annotations import json from collections.abc import Collection +from typing import TYPE_CHECKING from generalresearch.models.repdata.question import RepDataQuestion -from generalresearch.sql_helper import SqlHelper + +if TYPE_CHECKING: + from generalresearch.sql_helper import SqlHelper def get_profiling_library( diff --git a/generalresearch/managers/repdata/survey.py b/generalresearch/managers/repdata/survey.py index ab66374..2c3156c 100644 --- a/generalresearch/managers/repdata/survey.py +++ b/generalresearch/managers/repdata/survey.py @@ -3,6 +3,7 @@ from __future__ import annotations import json from collections.abc import Collection from datetime import UTC, datetime +from typing import TYPE_CHECKING import pymysql @@ -11,10 +12,12 @@ from generalresearch.managers.survey import SurveyManager from generalresearch.models.repdata.survey import ( RepDataCondition, RepDataStreamHashed, - RepDataSurvey, RepDataSurveyHashed, ) +if TYPE_CHECKING: + from generalresearch.models.repdata.survey import RepDataSurvey + SURVEY_FIELDS = [ "survey_id", "survey_uuid", diff --git a/generalresearch/managers/sago/profiling.py b/generalresearch/managers/sago/profiling.py index 4f5b2f5..3bbad3f 100644 --- a/generalresearch/managers/sago/profiling.py +++ b/generalresearch/managers/sago/profiling.py @@ -2,9 +2,12 @@ from __future__ import annotations import json from collections.abc import Collection +from typing import TYPE_CHECKING from generalresearch.models.sago.question import SagoQuestion -from generalresearch.sql_helper import SqlHelper + +if TYPE_CHECKING: + from generalresearch.sql_helper import SqlHelper def get_profiling_library( diff --git a/generalresearch/managers/spectrum/profiling.py b/generalresearch/managers/spectrum/profiling.py index 8a0904a..5de218c 100644 --- a/generalresearch/managers/spectrum/profiling.py +++ b/generalresearch/managers/spectrum/profiling.py @@ -2,9 +2,12 @@ from __future__ import annotations import json from collections.abc import Collection +from typing import TYPE_CHECKING from generalresearch.models.spectrum.question import SpectrumQuestion -from generalresearch.sql_helper import SqlHelper + +if TYPE_CHECKING: + from generalresearch.sql_helper import SqlHelper def get_profiling_library( diff --git a/generalresearch/managers/survey.py b/generalresearch/managers/survey.py index 1f057fb..65dbdd9 100644 --- a/generalresearch/managers/survey.py +++ b/generalresearch/managers/survey.py @@ -1,9 +1,12 @@ from __future__ import annotations from abc import ABC +from typing import TYPE_CHECKING from generalresearch.managers.base import SqlManager -from generalresearch.models.thl.survey import MarketplaceTask + +if TYPE_CHECKING: + from generalresearch.models.thl.survey import MarketplaceTask class SurveyManager(SqlManager, ABC): diff --git a/generalresearch/managers/thl/buyer.py b/generalresearch/managers/thl/buyer.py index 5aa2a01..1e20e2f 100644 --- a/generalresearch/managers/thl/buyer.py +++ b/generalresearch/managers/thl/buyer.py @@ -2,11 +2,14 @@ from __future__ import annotations from collections.abc import Collection from datetime import UTC, datetime +from typing import TYPE_CHECKING from generalresearch.managers.base import Permission, PostgresManager -from generalresearch.models import Source from generalresearch.models.thl.survey.buyer import Buyer -from generalresearch.pg_helper import PostgresConfig + +if TYPE_CHECKING: + from generalresearch.models import Source + from generalresearch.pg_helper import PostgresConfig class BuyerManager(PostgresManager): diff --git a/generalresearch/managers/thl/cashout_method.py b/generalresearch/managers/thl/cashout_method.py index f45e692..e701da3 100644 --- a/generalresearch/managers/thl/cashout_method.py +++ b/generalresearch/managers/thl/cashout_method.py @@ -3,20 +3,24 @@ from __future__ import annotations from collections.abc import Collection from copy import copy from datetime import UTC, datetime -from typing import Any +from typing import TYPE_CHECKING, Any from uuid import UUID, uuid4 from pydantic import NonNegativeInt from generalresearch.managers.base import PostgresManager -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, ) +if TYPE_CHECKING: + from generalresearch.models.thl.user import User + from generalresearch.models.thl.wallet.cashout_method import ( + CashMailCashoutMethodData, + PaypalCashoutMethodData, + ) + class CashoutMethodManager(PostgresManager): diff --git a/generalresearch/managers/thl/category.py b/generalresearch/managers/thl/category.py index e8a6aa6..e6a091b 100644 --- a/generalresearch/managers/thl/category.py +++ b/generalresearch/managers/thl/category.py @@ -1,11 +1,15 @@ from __future__ import annotations from collections.abc import Collection +from typing import TYPE_CHECKING -from generalresearch.managers.base import Permission, PostgresManager -from generalresearch.models.custom_types import UUIDStr +from generalresearch.managers.base import PostgresManager from generalresearch.models.thl.category import Category -from generalresearch.pg_helper import PostgresConfig + +if TYPE_CHECKING: + from generalresearch.managers.base import Permission + from generalresearch.models.custom_types import UUIDStr + from generalresearch.pg_helper import PostgresConfig class CategoryManager(PostgresManager): diff --git a/generalresearch/managers/thl/contest_manager.py b/generalresearch/managers/thl/contest_manager.py index 64206e1..3f85d31 100644 --- a/generalresearch/managers/thl/contest_manager.py +++ b/generalresearch/managers/thl/contest_manager.py @@ -2,7 +2,7 @@ from __future__ import annotations from collections.abc import Collection from datetime import UTC, datetime -from typing import Any, Literal, cast +from typing import TYPE_CHECKING, Any, Literal, cast from uuid import UUID import redis @@ -10,21 +10,10 @@ from pydantic import NonNegativeInt, PositiveInt from redis import Redis from generalresearch.managers.base import PostgresManager -from generalresearch.managers.thl.ledger_manager.thl_ledger import ( - ThlLedgerManager, -) -from generalresearch.managers.thl.user_manager.user_manager import ( - UserManager, -) -from generalresearch.models.custom_types import UUIDStr from generalresearch.models.thl.contest import ( ContestPrize, ContestWinner, ) -from generalresearch.models.thl.contest.contest import ( - Contest, - ContestUserView, -) from generalresearch.models.thl.contest.definitions import ( ContestStatus, ContestType, @@ -41,7 +30,6 @@ from generalresearch.models.thl.contest.leaderboard import ( LeaderboardContestUserView, ) from generalresearch.models.thl.contest.milestone import ( - ContestEntryTrigger, MilestoneContest, MilestoneEntry, MilestoneUserView, @@ -54,6 +42,20 @@ from generalresearch.models.thl.contest.raffle import ( ) from generalresearch.models.thl.user import User +if TYPE_CHECKING: + from generalresearch.managers.thl.ledger_manager.thl_ledger import ( + ThlLedgerManager, + ) + from generalresearch.managers.thl.user_manager.user_manager import ( + UserManager, + ) + from generalresearch.models.custom_types import UUIDStr + from generalresearch.models.thl.contest.contest import ( + Contest, + ContestUserView, + ) + from generalresearch.models.thl.contest.milestone import ContestEntryTrigger + CONTEST_SELECT = """ c.id, c.uuid::uuid, diff --git a/generalresearch/managers/thl/ipinfo.py b/generalresearch/managers/thl/ipinfo.py index 90757a2..93914c3 100644 --- a/generalresearch/managers/thl/ipinfo.py +++ b/generalresearch/managers/thl/ipinfo.py @@ -3,6 +3,7 @@ from __future__ import annotations import ipaddress from collections.abc import Collection from decimal import Decimal +from typing import TYPE_CHECKING import faker from grip_client.enums import AccessType @@ -14,17 +15,19 @@ from generalresearch.managers.base import ( PostgresManager, PostgresManagerWithRedis, ) -from generalresearch.models.custom_types import ( - CountryISOLike, - IPvAnyAddressStr, -) from generalresearch.models.thl.ipinfo import ( GeoIPInformation, IPGeoname, IPInformation, normalize_ip, ) -from generalresearch.pg_helper import PostgresConfig + +if TYPE_CHECKING: + from generalresearch.models.custom_types import ( + CountryISOLike, + IPvAnyAddressStr, + ) + from generalresearch.pg_helper import PostgresConfig fake = faker.Faker() diff --git a/generalresearch/managers/thl/ledger_manager/conditions.py b/generalresearch/managers/thl/ledger_manager/conditions.py index b2fd465..7dd3021 100644 --- a/generalresearch/managers/thl/ledger_manager/conditions.py +++ b/generalresearch/managers/thl/ledger_manager/conditions.py @@ -7,22 +7,23 @@ from typing import TYPE_CHECKING from generalresearch.config import JAMES_BILLINGS_BPID, JAMES_BILLINGS_TX_CUTOFF from generalresearch.currency import USDCent -from generalresearch.models.custom_types import UUIDStr -from generalresearch.models.thl.product import Product -from generalresearch.models.thl.session import Session, Wall -from generalresearch.models.thl.user import User - -logging.basicConfig() -logger = logging.getLogger("LedgerManager") -logger.setLevel(logging.INFO) 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.custom_types import UUIDStr + from generalresearch.models.thl.product import Product + from generalresearch.models.thl.session import Session, Wall + from generalresearch.models.thl.user import User + +logging.basicConfig() +logger = logging.getLogger("LedgerManager") +logger.setLevel(logging.INFO) def generate_condition_mp_payment(wall: Wall) -> Callable[..., bool]: diff --git a/generalresearch/managers/thl/ledger_manager/ledger.py b/generalresearch/managers/thl/ledger_manager/ledger.py index a5263f0..6cb4b28 100644 --- a/generalresearch/managers/thl/ledger_manager/ledger.py +++ b/generalresearch/managers/thl/ledger_manager/ledger.py @@ -4,7 +4,7 @@ import logging from collections import defaultdict from collections.abc import Callable, Collection from datetime import UTC, datetime, timedelta -from typing import Any +from typing import TYPE_CHECKING, Any from uuid import UUID import redis @@ -28,16 +28,19 @@ from generalresearch.managers.thl.ledger_manager.exceptions import ( LedgerTransactionFlagAlreadyExistsError, LedgerTransactionReleaseLockError, ) -from generalresearch.models.custom_types import UUIDStr, check_valid_uuid +from generalresearch.models.custom_types import check_valid_uuid from generalresearch.models.thl.ledger import ( LedgerAccount, LedgerEntry, LedgerTransaction, - UserLedgerTransactionType, UserLedgerTransactionTypesSummary, ) -from generalresearch.pg_helper import PostgresConfig -from generalresearch.redis_helper import RedisConfig + +if TYPE_CHECKING: + from generalresearch.models.custom_types import UUIDStr + from generalresearch.models.thl.ledger import UserLedgerTransactionType + from generalresearch.pg_helper import PostgresConfig + from generalresearch.redis_helper import RedisConfig logging.basicConfig() logger = logging.getLogger("LedgerManager") diff --git a/generalresearch/managers/thl/ledger_manager/thl_ledger.py b/generalresearch/managers/thl/ledger_manager/thl_ledger.py index ded6518..7aed619 100644 --- a/generalresearch/managers/thl/ledger_manager/thl_ledger.py +++ b/generalresearch/managers/thl/ledger_manager/thl_ledger.py @@ -29,15 +29,12 @@ from generalresearch.managers.thl.ledger_manager.conditions import ( from generalresearch.managers.thl.ledger_manager.ledger import ( LedgerManager, ) -from generalresearch.models.custom_types import UUIDStr -from generalresearch.models.thl.contest.contest import Contest from generalresearch.models.thl.contest.definitions import ( ContestPrizeKind, ContestType, ) from generalresearch.models.thl.contest.milestone import MilestoneContest from generalresearch.models.thl.contest.raffle import ( - ContestEntry, ContestEntryType, RaffleContest, ) @@ -53,14 +50,20 @@ from generalresearch.models.thl.ledger import ( from generalresearch.models.thl.ledger import ( TransactionMetadataColumns as tmc, ) -from generalresearch.models.thl.payout import UserPayoutEvent from generalresearch.models.thl.product import Product -from generalresearch.models.thl.session import Session, Status, Wall -from generalresearch.models.thl.user import User +from generalresearch.models.thl.session import Status from generalresearch.models.thl.wallet import PayoutType if TYPE_CHECKING: - from generalresearch.models.thl.contest.contest import ContestWinner + from generalresearch.models.custom_types import UUIDStr + from generalresearch.models.thl.contest.contest import Contest, ContestWinner + from generalresearch.models.thl.contest.raffle import ( + ContestEntry, + ) + from generalresearch.models.thl.ledger import LedgerTransaction + from generalresearch.models.thl.payout import UserPayoutEvent + from generalresearch.models.thl.session import Session, Wall + from generalresearch.models.thl.user import User logging.basicConfig() logger = logging.getLogger("LedgerManager") @@ -663,9 +666,7 @@ class ThlLedgerManager(LedgerManager): ) else: - logger.info( - "create_transaction_bp_adjustment. No transactions needed." - ) + logger.info("create_transaction_bp_adjustment. No transactions needed.") return None else: new_bp_payout = new_payout @@ -735,9 +736,7 @@ class ThlLedgerManager(LedgerManager): ) else: - logger.info( - "create_transaction_bp_adjustment. No transactions needed." - ) + logger.info("create_transaction_bp_adjustment. No transactions needed.") return None logger.info(entries) @@ -796,9 +795,7 @@ class ThlLedgerManager(LedgerManager): if skip_one_per_day_check or skip_wallet_balance_check: skip_flag_check = True - assert ( - datetime.now(tz=UTC) > created - ), "created cannot be in the future" + assert datetime.now(tz=UTC) > created, "created cannot be in the future" f = lambda: self.create_tx_bp_payout_( product=product, amount=amount, @@ -902,9 +899,7 @@ class ThlLedgerManager(LedgerManager): :param skip_flag_check: If True, we skip the flag check to allow for retry of a failed previous call. """ - assert ( - datetime.now(tz=UTC) > created - ), "created cannot be in the future" + assert datetime.now(tz=UTC) > created, "created cannot be in the future" assert isinstance(amount, int) assert isinstance(amount, USDCent) diff --git a/generalresearch/managers/thl/payout.py b/generalresearch/managers/thl/payout.py index f50e0d2..2914ba4 100644 --- a/generalresearch/managers/thl/payout.py +++ b/generalresearch/managers/thl/payout.py @@ -2,7 +2,7 @@ from __future__ import annotations from collections.abc import Collection from datetime import UTC, datetime -from typing import Any +from typing import TYPE_CHECKING, Any from uuid import uuid4 import numpy as np @@ -19,16 +19,9 @@ from generalresearch.managers.base import ( from generalresearch.managers.thl.ledger_manager.exceptions import ( LedgerTransactionConditionFailedError, ) -from generalresearch.managers.thl.ledger_manager.thl_ledger import ( - ThlLedgerManager, -) -from generalresearch.managers.thl.product import ProductManager -from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr -from generalresearch.models.gr.business import Business from generalresearch.models.thl.definitions import PayoutStatus from generalresearch.models.thl.ledger import ( Direction, - LedgerAccount, OrderBy, ) from generalresearch.models.thl.payout import ( @@ -38,13 +31,22 @@ from generalresearch.models.thl.payout import ( PayoutEvent, UserPayoutEvent, ) -from generalresearch.models.thl.product import Product from generalresearch.models.thl.wallet import PayoutType from generalresearch.models.thl.wallet.cashout_method import ( CashMailOrderData, CashoutRequestInfo, ) +if TYPE_CHECKING: + from generalresearch.managers.thl.ledger_manager.thl_ledger import ( + ThlLedgerManager, + ) + from generalresearch.managers.thl.product import ProductManager + from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr + from generalresearch.models.gr.business import Business + from generalresearch.models.thl.ledger import LedgerAccount + from generalresearch.models.thl.product import Product + class PayoutEventManager(PostgresManagerWithRedis): """This is the default base Payout Event Manger. It acts as a base for diff --git a/generalresearch/managers/thl/product.py b/generalresearch/managers/thl/product.py index aac2979..535e566 100644 --- a/generalresearch/managers/thl/product.py +++ b/generalresearch/managers/thl/product.py @@ -18,15 +18,15 @@ from sentry_sdk import capture_exception from generalresearch.decorators import LOG from generalresearch.managers.base import ( - Permission, PostgresManager, ) -from generalresearch.models.custom_types import UUIDStr, is_valid_uuid -from generalresearch.pg_helper import PostgresConfig - -logger = logging.getLogger() +from generalresearch.models.custom_types import is_valid_uuid if TYPE_CHECKING: + from generalresearch.managers.base import ( + Permission, + ) + from generalresearch.models.custom_types import UUIDStr from generalresearch.models.thl.product import ( PayoutConfig, Product, @@ -38,6 +38,9 @@ if TYPE_CHECKING: UserHealthConfig, UserWalletConfig, ) + from generalresearch.pg_helper import PostgresConfig + +logger = logging.getLogger() class ProductManager(PostgresManager): diff --git a/generalresearch/managers/thl/profiling/uqa.py b/generalresearch/managers/thl/profiling/uqa.py index fa1747b..d240eea 100644 --- a/generalresearch/managers/thl/profiling/uqa.py +++ b/generalresearch/managers/thl/profiling/uqa.py @@ -3,13 +3,16 @@ from __future__ import annotations import logging from collections.abc import Collection from datetime import UTC, datetime, timedelta +from typing import TYPE_CHECKING from generalresearch.managers.base import PostgresManagerWithRedis from generalresearch.models.thl.profiling.user_question_answer import ( DUMMY_UQA, UserQuestionAnswer, ) -from generalresearch.models.thl.user import User + +if TYPE_CHECKING: + from generalresearch.models.thl.user import User logger = logging.getLogger() diff --git a/generalresearch/managers/thl/profiling/user_upk.py b/generalresearch/managers/thl/profiling/user_upk.py index 6106037..53475b5 100644 --- a/generalresearch/managers/thl/profiling/user_upk.py +++ b/generalresearch/managers/thl/profiling/user_upk.py @@ -4,27 +4,29 @@ import json from collections import defaultdict from collections.abc import Collection from datetime import UTC, datetime, timedelta -from typing import Any +from typing import TYPE_CHECKING, Any from uuid import UUID from psycopg import Cursor from pydantic import PositiveInt from generalresearch.managers.base import ( - Permission, PostgresManagerWithRedis, ) from generalresearch.managers.thl.profiling.schema import UpkSchemaManager from generalresearch.models.thl.profiling.upk_property import ( Cardinality, PropertyType, - UpkProperty, ) from generalresearch.models.thl.profiling.upk_question_answer import ( UpkQuestionAnswer, ) -from generalresearch.pg_helper import PostgresConfig -from generalresearch.redis_helper import RedisConfig + +if TYPE_CHECKING: + from generalresearch.managers.base import Permission + from generalresearch.models.thl.profiling.upk_property import UpkProperty + from generalresearch.pg_helper import PostgresConfig + from generalresearch.redis_helper import RedisConfig class UserUpkManager(PostgresManagerWithRedis): diff --git a/generalresearch/managers/thl/session.py b/generalresearch/managers/thl/session.py index 959003c..7f17252 100644 --- a/generalresearch/managers/thl/session.py +++ b/generalresearch/managers/thl/session.py @@ -3,7 +3,7 @@ from __future__ import annotations from collections.abc import Collection from datetime import UTC, datetime, timedelta from decimal import Decimal -from typing import Any +from typing import TYPE_CHECKING, Any from uuid import UUID, uuid4 from faker import Faker @@ -16,14 +16,7 @@ from generalresearch.managers.base import ( PostgresManager, ) from generalresearch.managers.thl.product import ProductManager -from generalresearch.models import DeviceType -from generalresearch.models.custom_types import UUIDStr from generalresearch.models.legacy.bucket import Bucket -from generalresearch.models.thl.definitions import ( - SessionStatusCode2, - Status, - StatusCode1, -) from generalresearch.models.thl.session import ( Session, Wall, @@ -34,6 +27,15 @@ 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.thl.definitions import ( + SessionStatusCode2, + Status, + StatusCode1, + ) + fake = Faker() @@ -190,7 +192,12 @@ class SessionManager(PostgresManager): # re-run model_validate after finished = finished if finished else datetime.now(tz=UTC) session.update( - status=status, status_code_1=status_code_1, status_code_2=status_code_2, finished=finished, payout=payout, user_payout=user_payout + status=status, + status_code_1=status_code_1, + status_code_2=status_code_2, + finished=finished, + payout=payout, + user_payout=user_payout, ) d = session.model_dump_mysql() self.pg_config.execute_write( @@ -446,9 +453,7 @@ class SessionManager(PostgresManager): if started_before or started_after: started_after = started_after or datetime(2017, 1, 1, tzinfo=UTC) started_before = started_before or datetime.now(tz=UTC) - assert ( - started_after.tzinfo == UTC - ), "started_after must be tz-aware as UTC" + assert started_after.tzinfo == UTC, "started_after must be tz-aware as UTC" assert ( started_before.tzinfo == UTC ), "started_before must be tz-aware as UTC" diff --git a/generalresearch/managers/thl/survey.py b/generalresearch/managers/thl/survey.py index 024ad38..eacb345 100644 --- a/generalresearch/managers/thl/survey.py +++ b/generalresearch/managers/thl/survey.py @@ -3,7 +3,7 @@ from __future__ import annotations from collections import defaultdict from collections.abc import Collection from datetime import UTC, datetime -from typing import Any +from typing import TYPE_CHECKING, Any import pandas as pd from more_itertools import chunked @@ -14,12 +14,14 @@ 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.custom_types import SurveyKey from generalresearch.models.thl.survey.model import ( Survey, SurveyStat, ) -from generalresearch.pg_helper import PostgresConfig + +if TYPE_CHECKING: + from generalresearch.models.custom_types import SurveyKey + from generalresearch.pg_helper import PostgresConfig class SurveyManager(PostgresManager): diff --git a/generalresearch/managers/thl/survey_penalty.py b/generalresearch/managers/thl/survey_penalty.py index bf914cb..08f8649 100644 --- a/generalresearch/managers/thl/survey_penalty.py +++ b/generalresearch/managers/thl/survey_penalty.py @@ -4,21 +4,24 @@ import json import threading from collections import defaultdict from datetime import timedelta +from typing import TYPE_CHECKING from cachetools import TTLCache, cachedmethod from generalresearch.decorators import LOG from generalresearch.managers.base import RedisManager -from generalresearch.models.custom_types import ( - UUIDStr, -) -from generalresearch.models.thl.survey.penalty import ( - BPSurveyPenalty, - Penalty, - PenaltyListAdapter, - TeamSurveyPenalty, -) -from generalresearch.redis_helper import RedisConfig +from generalresearch.models.thl.survey.penalty import PenaltyListAdapter + +if TYPE_CHECKING: + from generalresearch.models.custom_types import ( + UUIDStr, + ) + from generalresearch.models.thl.survey.penalty import ( + BPSurveyPenalty, + Penalty, + TeamSurveyPenalty, + ) + from generalresearch.redis_helper import RedisConfig class SurveyPenaltyManager(RedisManager): diff --git a/generalresearch/managers/thl/task_adjustment.py b/generalresearch/managers/thl/task_adjustment.py index 2d89334..e3f382d 100644 --- a/generalresearch/managers/thl/task_adjustment.py +++ b/generalresearch/managers/thl/task_adjustment.py @@ -4,17 +4,14 @@ import logging from datetime import UTC, datetime 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.ledger_manager.thl_ledger import ( - ThlLedgerManager, -) from generalresearch.managers.thl.session import SessionManager from generalresearch.managers.thl.wall import WallManager -from generalresearch.models.custom_types import UUIDStr from generalresearch.models.thl.definitions import ( Status, WallAdjustedStatus, @@ -24,6 +21,12 @@ from generalresearch.models.thl.session import ( ) from generalresearch.models.thl.task_adjustment import TaskAdjustmentEvent +if TYPE_CHECKING: + from generalresearch.managers.thl.ledger_manager.thl_ledger import ( + ThlLedgerManager, + ) + from generalresearch.models.custom_types import UUIDStr + logging.basicConfig() logger = logging.getLogger(__name__) diff --git a/generalresearch/managers/thl/user_compensate.py b/generalresearch/managers/thl/user_compensate.py index 8338424..3f0f3ec 100644 --- a/generalresearch/managers/thl/user_compensate.py +++ b/generalresearch/managers/thl/user_compensate.py @@ -2,15 +2,17 @@ from __future__ import annotations from datetime import UTC, datetime from decimal import Decimal +from typing import TYPE_CHECKING from uuid import uuid4 from pydantic import NonNegativeInt -from generalresearch.managers.thl.ledger_manager.thl_ledger import ( - ThlLedgerManager, -) -from generalresearch.models.custom_types import UUIDStr -from generalresearch.models.thl.user import User +if TYPE_CHECKING: + from generalresearch.managers.thl.ledger_manager.thl_ledger import ( + ThlLedgerManager, + ) + from generalresearch.models.custom_types import UUIDStr + from generalresearch.models.thl.user import User def user_compensate( diff --git a/generalresearch/managers/thl/user_manager/__init__.py b/generalresearch/managers/thl/user_manager/__init__.py index 449dc2a..6414ef2 100644 --- a/generalresearch/managers/thl/user_manager/__init__.py +++ b/generalresearch/managers/thl/user_manager/__init__.py @@ -4,11 +4,12 @@ import csv import logging from pathlib import Path from threading import RLock -from typing import Any +from typing import TYPE_CHECKING, Any from cachetools import TTLCache, cached -from generalresearch.models.thl.product import Product +if TYPE_CHECKING: + from generalresearch.models.thl.product import Product logger = logging.getLogger() diff --git a/generalresearch/managers/thl/user_manager/mysql_user_manager.py b/generalresearch/managers/thl/user_manager/mysql_user_manager.py index af65d65..2af23dc 100644 --- a/generalresearch/managers/thl/user_manager/mysql_user_manager.py +++ b/generalresearch/managers/thl/user_manager/mysql_user_manager.py @@ -4,14 +4,17 @@ import logging from collections.abc import Collection from datetime import UTC, datetime from functools import lru_cache +from typing import TYPE_CHECKING from uuid import uuid4 import psycopg from psycopg import sql -from generalresearch.models.custom_types import UUIDStr from generalresearch.models.thl.user import User -from generalresearch.pg_helper import PostgresConfig + +if TYPE_CHECKING: + from generalresearch.models.custom_types import UUIDStr + from generalresearch.pg_helper import PostgresConfig logging.basicConfig() logger = logging.getLogger() diff --git a/generalresearch/managers/thl/user_manager/rate_limit.py b/generalresearch/managers/thl/user_manager/rate_limit.py index 5d239c6..2aaa134 100644 --- a/generalresearch/managers/thl/user_manager/rate_limit.py +++ b/generalresearch/managers/thl/user_manager/rate_limit.py @@ -1,4 +1,5 @@ import logging +from typing import TYPE_CHECKING from limits import RateLimitItem, RateLimitItemPerHour, storage, strategies from limits.limits import TIME_TYPES, safe_string @@ -10,7 +11,9 @@ from generalresearch.managers.thl.user_manager import ( from generalresearch.managers.thl.user_manager.exceptions import ( UserCreateNotAllowedError, ) -from generalresearch.models.thl.product import Product + +if TYPE_CHECKING: + from generalresearch.models.thl.product import Product logger = logging.getLogger() diff --git a/generalresearch/managers/thl/user_manager/user_manager.py b/generalresearch/managers/thl/user_manager/user_manager.py index 26b1fd6..907a030 100644 --- a/generalresearch/managers/thl/user_manager/user_manager.py +++ b/generalresearch/managers/thl/user_manager/user_manager.py @@ -10,7 +10,9 @@ from pydantic import RedisDsn from generalresearch.managers.base import Permission from generalresearch.managers.thl.product import ProductManager -from generalresearch.managers.thl.user_manager.exceptions import UserDoesntExistError +from generalresearch.managers.thl.user_manager.exceptions import ( + UserDoesntExistError, +) from generalresearch.managers.thl.user_manager.mysql_user_manager import ( MysqlUserManager, ) @@ -20,15 +22,16 @@ from generalresearch.managers.thl.user_manager.rate_limit import ( from generalresearch.managers.thl.user_manager.redis_user_manager import ( RedisUserManager, ) -from generalresearch.pg_helper import PostgresConfig from generalresearch.utils.copying_cache import deepcopy_return if TYPE_CHECKING: + from generalresearch.managers.thl.userhealth import AuditLogManager from generalresearch.models.custom_types import UUIDStr from generalresearch.models.thl.product import Product from generalresearch.models.thl.user import User from generalresearch.models.thl.userhealth import AuditLog + from generalresearch.pg_helper import PostgresConfig logging.basicConfig() logger = logging.getLogger() diff --git a/generalresearch/managers/thl/userhealth.py b/generalresearch/managers/thl/userhealth.py index 2bbdfab..b986256 100644 --- a/generalresearch/managers/thl/userhealth.py +++ b/generalresearch/managers/thl/userhealth.py @@ -11,24 +11,25 @@ from pydantic import NonNegativeInt, PositiveInt from generalresearch.decorators import LOG from generalresearch.managers.base import ( - Permission, PostgresManager, PostgresManagerWithRedis, ) from generalresearch.managers.thl.ipinfo import GeoIpInfoManager -from generalresearch.models.custom_types import IPvAnyAddressStr -from generalresearch.models.thl.product import Product from generalresearch.models.thl.user_iphistory import ( IPRecord, UserIPHistory, UserIPRecord, ) -from generalresearch.models.thl.userhealth import AuditLog, AuditLogLevel -from generalresearch.pg_helper import PostgresConfig -from generalresearch.redis_helper import RedisConfig +from generalresearch.models.thl.userhealth import AuditLog if TYPE_CHECKING: + from generalresearch.managers.base import Permission + from generalresearch.models.custom_types import IPvAnyAddressStr + from generalresearch.models.thl.product import Product from generalresearch.models.thl.user import User + from generalresearch.models.thl.userhealth import AuditLogLevel + from generalresearch.pg_helper import PostgresConfig + from generalresearch.redis_helper import RedisConfig fake = faker.Faker() diff --git a/generalresearch/managers/thl/wall.py b/generalresearch/managers/thl/wall.py index 774db9e..b9dc94d 100644 --- a/generalresearch/managers/thl/wall.py +++ b/generalresearch/managers/thl/wall.py @@ -6,6 +6,7 @@ from collections.abc import Collection from datetime import UTC, datetime, timedelta from decimal import Decimal from functools import cached_property +from typing import TYPE_CHECKING from uuid import uuid4 from faker import Faker @@ -15,18 +16,12 @@ from pydantic import AwareDatetime, PositiveInt from generalresearch.managers import parse_order_by from generalresearch.managers.base import ( - Permission, PostgresManager, PostgresManagerWithRedis, ) from generalresearch.models import Source -from generalresearch.models.custom_types import SurveyKey, UUIDStr from generalresearch.models.thl.definitions import ( - ReportValue, - Status, - StatusCode1, WallAdjustedStatus, - WallStatusCode2, ) from generalresearch.models.thl.ledger import OrderBy from generalresearch.models.thl.session import ( @@ -35,7 +30,19 @@ from generalresearch.models.thl.session import ( check_adjusted_status_wall_consistent, ) from generalresearch.models.thl.survey.model import TaskActivity -from generalresearch.pg_helper import PostgresConfig + +if TYPE_CHECKING: + from generalresearch.managers.base import ( + Permission, + ) + from generalresearch.models.custom_types import SurveyKey, UUIDStr + from generalresearch.models.thl.definitions import ( + ReportValue, + Status, + StatusCode1, + WallStatusCode2, + ) + from generalresearch.pg_helper import PostgresConfig logger = logging.getLogger("WallManager") fake = Faker() diff --git a/generalresearch/managers/thl/wallet/__init__.py b/generalresearch/managers/thl/wallet/__init__.py index 05e700d..457483f 100644 --- a/generalresearch/managers/thl/wallet/__init__.py +++ b/generalresearch/managers/thl/wallet/__init__.py @@ -1,28 +1,30 @@ from decimal import Decimal -from typing import Any +from typing import TYPE_CHECKING, Any -from generalresearch.managers.thl.ledger_manager.thl_ledger import ( - ThlLedgerManager, -) -from generalresearch.managers.thl.payout import ( - PayoutEventManager, - UserPayoutEventManager, -) -from generalresearch.managers.thl.tango_api import TangoClient -from generalresearch.managers.thl.user_manager.user_manager import ( - UserManager, -) -from generalresearch.managers.thl.userhealth import UserIpHistoryManager from generalresearch.managers.thl.wallet.approve import ( approve_amt_cashout, approve_paypal_order, ) from generalresearch.models.thl.definitions import PayoutStatus -from generalresearch.models.thl.payout import UserPayoutEvent from generalresearch.models.thl.wallet import PayoutType -from generalresearch.models.thl.wallet.cashout_method import ( - CashMailOrderData, -) + +if TYPE_CHECKING: + from generalresearch.managers.thl.ledger_manager.thl_ledger import ( + ThlLedgerManager, + ) + from generalresearch.managers.thl.payout import ( + PayoutEventManager, + UserPayoutEventManager, + ) + from generalresearch.managers.thl.tango_api import TangoClient + from generalresearch.managers.thl.user_manager.user_manager import ( + UserManager, + ) + from generalresearch.managers.thl.userhealth import UserIpHistoryManager + from generalresearch.models.thl.payout import UserPayoutEvent + from generalresearch.models.thl.wallet.cashout_method import ( + CashMailOrderData, + ) def manage_pending_cashout( diff --git a/generalresearch/managers/thl/wallet/approve.py b/generalresearch/managers/thl/wallet/approve.py index 4a7ae5e..7fedec1 100644 --- a/generalresearch/managers/thl/wallet/approve.py +++ b/generalresearch/managers/thl/wallet/approve.py @@ -1,10 +1,16 @@ -from generalresearch.managers.thl.ledger_manager.thl_ledger import ( - ThlLedgerManager, -) -from generalresearch.managers.thl.payout import PayoutEventManager +from __future__ import annotations + +from typing import TYPE_CHECKING + from generalresearch.models.thl.definitions import PayoutStatus -from generalresearch.models.thl.payout import UserPayoutEvent -from generalresearch.models.thl.user import User + +if TYPE_CHECKING: + from generalresearch.managers.thl.ledger_manager.thl_ledger import ( + ThlLedgerManager, + ) + from generalresearch.managers.thl.payout import PayoutEventManager + from generalresearch.models.thl.payout import UserPayoutEvent + from generalresearch.models.thl.user import User def approve_paypal_order( diff --git a/generalresearch/managers/thl/wallet/tango.py b/generalresearch/managers/thl/wallet/tango.py index be8fd97..038f67f 100644 --- a/generalresearch/managers/thl/wallet/tango.py +++ b/generalresearch/managers/thl/wallet/tango.py @@ -1,16 +1,21 @@ -from typing import Any +from __future__ import annotations + +from typing import TYPE_CHECKING, Any from generalresearch.config import ( is_debug, ) -from generalresearch.managers.thl.ledger_manager.thl_ledger import ( - ThlLedgerManager, -) -from generalresearch.managers.thl.payout import PayoutEventManager -from generalresearch.managers.thl.tango_api import TangoClient, TangoOrderRequest +from generalresearch.managers.thl.tango_api import TangoOrderRequest from generalresearch.models.thl.definitions import PayoutStatus -from generalresearch.models.thl.payout import UserPayoutEvent -from generalresearch.models.thl.user import User + +if TYPE_CHECKING: + from generalresearch.managers.thl.ledger_manager.thl_ledger import ( + ThlLedgerManager, + ) + from generalresearch.managers.thl.payout import PayoutEventManager + from generalresearch.managers.thl.tango_api import TangoClient + from generalresearch.models.thl.payout import UserPayoutEvent + from generalresearch.models.thl.user import User def complete_tango_order( diff --git a/generalresearch/models/admin/request.py b/generalresearch/models/admin/request.py index 6112786..f128e1b 100644 --- a/generalresearch/models/admin/request.py +++ b/generalresearch/models/admin/request.py @@ -2,12 +2,13 @@ from __future__ import annotations from datetime import UTC, datetime, timedelta from enum import Enum -from typing import Literal +from typing import TYPE_CHECKING, Literal import pandas as pd from pydantic import BaseModel, Field, computed_field, model_validator -from generalresearch.models.custom_types import AwareDatetimeISO +if TYPE_CHECKING: + from generalresearch.models.custom_types import AwareDatetimeISO class ReportType(Enum): diff --git a/generalresearch/models/cint/question.py b/generalresearch/models/cint/question.py index 4c0e52f..44efd13 100644 --- a/generalresearch/models/cint/question.py +++ b/generalresearch/models/cint/question.py @@ -9,14 +9,14 @@ from uuid import UUID from pydantic import BaseModel, Field, field_validator, model_validator from generalresearch.models import Source, string_utils -from generalresearch.models.cint import CintQuestionIdType -from generalresearch.models.custom_types import AwareDatetimeISO from generalresearch.models.thl.profiling.marketplace import ( MarketplaceQuestion, MarketplaceUserQuestionAnswer, ) if TYPE_CHECKING: + from generalresearch.models.cint import CintQuestionIdType + from generalresearch.models.custom_types import AwareDatetimeISO from generalresearch.models.thl.profiling.upk_question import ( UpkQuestion, ) diff --git a/generalresearch/models/cint/survey.py b/generalresearch/models/cint/survey.py index fde4559..8c8f882 100644 --- a/generalresearch/models/cint/survey.py +++ b/generalresearch/models/cint/survey.py @@ -4,7 +4,7 @@ import json import logging from datetime import UTC, datetime from decimal import Decimal -from typing import Annotated, Any, Literal, Self +from typing import TYPE_CHECKING, Annotated, Any, Literal, Self from more_itertools import flatten from pydantic import ( @@ -19,12 +19,6 @@ from pydantic import ( from generalresearch.locales import Localelator from generalresearch.models import Source, TaskCalculationType -from generalresearch.models.cint import CintQuestionIdType -from generalresearch.models.custom_types import ( - AlphaNumStr, - AwareDatetimeISO, - CoercedStr, -) from generalresearch.models.thl.demographics import Gender from generalresearch.models.thl.survey import MarketplaceTask from generalresearch.models.thl.survey.condition import ( @@ -32,6 +26,14 @@ from generalresearch.models.thl.survey.condition import ( MarketplaceCondition, ) +if TYPE_CHECKING: + from generalresearch.models.cint import CintQuestionIdType + from generalresearch.models.custom_types import ( + AlphaNumStr, + AwareDatetimeISO, + CoercedStr, + ) + logging.basicConfig() logger = logging.getLogger() logger.setLevel(logging.INFO) diff --git a/generalresearch/models/cint/task_collection.py b/generalresearch/models/cint/task_collection.py index 4ae8de4..31a0173 100644 --- a/generalresearch/models/cint/task_collection.py +++ b/generalresearch/models/cint/task_collection.py @@ -1,15 +1,19 @@ from __future__ import annotations +from typing import TYPE_CHECKING + import pandas as pd from pandera.pandas import Check, Column, DataFrameSchema, Index from generalresearch.locales import Localelator -from generalresearch.models.cint.survey import CintSurvey from generalresearch.models.thl.survey.task_collection import ( TaskCollection, create_empty_df_from_schema, ) +if TYPE_CHECKING: + from generalresearch.models.cint.survey import CintSurvey + COUNTRY_ISOS: set[str] = Localelator().get_all_countries() LANGUAGE_ISOS: set[str] = Localelator().get_all_languages() diff --git a/generalresearch/models/dynata/question.py b/generalresearch/models/dynata/question.py index b95f7c6..1ed560a 100644 --- a/generalresearch/models/dynata/question.py +++ b/generalresearch/models/dynata/question.py @@ -7,14 +7,16 @@ import re from datetime import timedelta from enum import StrEnum from functools import cached_property -from typing import Any, Literal +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.custom_types import AwareDatetimeISO from generalresearch.models.thl.profiling.marketplace import MarketplaceQuestion +if TYPE_CHECKING: + from generalresearch.models.custom_types import AwareDatetimeISO + logging.basicConfig() logger = logging.getLogger() logger.setLevel(logging.INFO) diff --git a/generalresearch/models/dynata/survey.py b/generalresearch/models/dynata/survey.py index 942ab4f..70e3659 100644 --- a/generalresearch/models/dynata/survey.py +++ b/generalresearch/models/dynata/survey.py @@ -5,7 +5,7 @@ import logging from datetime import UTC from decimal import Decimal from functools import cached_property -from typing import Any, Literal, Self +from typing import TYPE_CHECKING, Any, Literal, Self from more_itertools import flatten from pydantic import ( @@ -19,14 +19,7 @@ from pydantic import ( ) from generalresearch.locales import Localelator -from generalresearch.models import Source, TaskCalculationType -from generalresearch.models.custom_types import ( - AlphaNumStr, - AlphaNumStrSet, - AwareDatetimeISO, - CoercedStr, - DeviceTypes, -) +from generalresearch.models import Source from generalresearch.models.dynata import DynataStatus from generalresearch.models.thl.demographics import ( Gender, @@ -37,6 +30,16 @@ from generalresearch.models.thl.survey.condition import ( MarketplaceCondition, ) +if TYPE_CHECKING: + from generalresearch.models import TaskCalculationType + from generalresearch.models.custom_types import ( + AlphaNumStr, + AlphaNumStrSet, + AwareDatetimeISO, + CoercedStr, + DeviceTypes, + ) + logging.basicConfig() logger = logging.getLogger() logger.setLevel(logging.INFO) diff --git a/generalresearch/models/dynata/task_collection.py b/generalresearch/models/dynata/task_collection.py index 2b82bfd..94868bb 100644 --- a/generalresearch/models/dynata/task_collection.py +++ b/generalresearch/models/dynata/task_collection.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import Any +from typing import TYPE_CHECKING, Any import pandas as pd from pandera.pandas import Check, Column, DataFrameSchema, Index @@ -8,12 +8,14 @@ from pandera.pandas import Check, Column, DataFrameSchema, Index from generalresearch.locales import Localelator from generalresearch.models import TaskCalculationType from generalresearch.models.dynata import DynataStatus -from generalresearch.models.dynata.survey import DynataSurvey from generalresearch.models.thl.survey.task_collection import ( TaskCollection, create_empty_df_from_schema, ) +if TYPE_CHECKING: + from generalresearch.models.dynata.survey import DynataSurvey + COUNTRY_ISOS = Localelator().get_all_countries() LANGUAGE_ISOS = Localelator().get_all_languages() diff --git a/generalresearch/models/events.py b/generalresearch/models/events.py index 70f699b..8d059f9 100644 --- a/generalresearch/models/events.py +++ b/generalresearch/models/events.py @@ -1,6 +1,6 @@ from datetime import UTC, datetime, timedelta from enum import StrEnum -from typing import Annotated, Literal +from typing import TYPE_CHECKING, Annotated, Literal from uuid import uuid4 from pydantic import ( @@ -13,18 +13,19 @@ from pydantic import ( model_validator, ) -from generalresearch.models import Source -from generalresearch.models.custom_types import ( - AwareDatetimeISO, - CountryISOLike, - UUIDStr, -) -from generalresearch.models.thl.definitions import ( - SessionStatusCode2, - Status, - StatusCode1, - WallStatusCode2, -) +if TYPE_CHECKING: + from generalresearch.models import Source + from generalresearch.models.custom_types import ( + AwareDatetimeISO, + CountryISOLike, + UUIDStr, + ) + from generalresearch.models.thl.definitions import ( + SessionStatusCode2, + Status, + StatusCode1, + WallStatusCode2, + ) class MessageKind(StrEnum): diff --git a/generalresearch/models/gr/authentication.py b/generalresearch/models/gr/authentication.py index 25f65fa..21ece8c 100644 --- a/generalresearch/models/gr/authentication.py +++ b/generalresearch/models/gr/authentication.py @@ -17,14 +17,14 @@ from pydantic import ( ) from generalresearch.decorators import LOG -from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr -from generalresearch.pg_helper import PostgresConfig -from generalresearch.redis_helper import RedisConfig if TYPE_CHECKING: + from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr 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 + from generalresearch.redis_helper import RedisConfig class Claims(BaseModel): diff --git a/generalresearch/models/gr/business.py b/generalresearch/models/gr/business.py index f6689a8..146d690 100644 --- a/generalresearch/models/gr/business.py +++ b/generalresearch/models/gr/business.py @@ -20,24 +20,28 @@ from pydantic_extra_types.phone_numbers import PhoneNumber from generalresearch.currency import USDCent from generalresearch.decorators import LOG -from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge 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 LedgerAccount, OrderBy -from generalresearch.models.thl.payout import BusinessPayoutEvent -from generalresearch.pg_helper import PostgresConfig -from generalresearch.redis_helper import RedisConfig +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 + 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 + from generalresearch.redis_helper import RedisConfig + logging.basicConfig() logger = logging.getLogger(__name__) @@ -49,6 +53,9 @@ if TYPE_CHECKING: from generalresearch.incite.mergers.foundations.enriched_wall import ( EnrichedWallMerge, ) + from generalresearch.managers.gr.business import ( + BusinessBankAccountManager, + ) from generalresearch.managers.thl.ledger_manager.ledger import ( LedgerManager, ) @@ -58,7 +65,7 @@ if TYPE_CHECKING: from generalresearch.managers.thl.payout import ( BusinessPayoutEventManager, ) - from generalresearch.models.gr.team import Team + from generalresearch.managers.thl.product import ProductManager from generalresearch.models.thl.product import Product @@ -263,8 +270,6 @@ class Business(BaseModel): self.addresses = [BusinessAddress.model_validate(i) for i in res] def prefetch_teams(self, pg_config: PostgresConfig) -> None: - from generalresearch.models.gr.team import Team - with pg_config.make_connection() as conn, conn.cursor( row_factory=dict_row ) as c: @@ -288,30 +293,27 @@ class Business(BaseModel): self.teams = [Team.model_validate(i) for i in res] - def prefetch_products(self, thl_pg_config: PostgresConfig) -> None: + def prefetch_products(self, product_manager: ProductManager) -> None: """ :return: All the Products for this Business """ - from generalresearch.managers.thl.product import ProductManager - pm = ProductManager(pg_config=thl_pg_config) - self.products = pm.fetch_uuids(business_uuids=[self.uuid]) + self.products = product_manager.fetch_uuids(business_uuids=[self.uuid]) - def prefetch_bank_accounts(self, pg_config: PostgresConfig) -> None: - from generalresearch.managers.gr.business import ( - BusinessBankAccountManager, + def prefetch_bank_accounts( + self, business_bank_account_manager: BusinessBankAccountManager + ) -> None: + self.bank_accounts = business_bank_account_manager.get_by_business_id( + business_id=self.id ) - bam = BusinessBankAccountManager(pg_config=pg_config) - self.bank_accounts = bam.get_by_business_id(business_id=self.id) - def prefetch_bp_accounts( - self, thl_lm: ThlLedgerManager, thl_pg_config: PostgresConfig + self, thl_lm: ThlLedgerManager, product_manager: ProductManager ): # We need to prefetch the Products everytime because there is no way # of knowing if a new Product has been added since the last time it # ran. - self.prefetch_products(thl_pg_config=thl_pg_config) + self.prefetch_products(product_manager=product_manager) product_lookup = {p.uuid: p for p in self.products} accounts = thl_lm.get_accounts_if_exists( @@ -332,6 +334,7 @@ class Business(BaseModel): ) product = product_lookup[product_uuid] thl_lm.get_account_or_create_bp_wallet(product=product) + if refresh: accounts = thl_lm.get_accounts_if_exists( qualified_names=[ diff --git a/generalresearch/models/gr/team.py b/generalresearch/models/gr/team.py index 4752bea..b36ac4c 100644 --- a/generalresearch/models/gr/team.py +++ b/generalresearch/models/gr/team.py @@ -22,27 +22,33 @@ from pydantic import ( from pydantic.json_schema import SkipJsonSchema from generalresearch.decorators import LOG -from generalresearch.incite.mergers.foundations.enriched_session import ( - EnrichedSessionMerge, -) -from generalresearch.incite.mergers.foundations.enriched_wall import ( - EnrichedWallMerge, -) from generalresearch.models.admin.request import ReportRequest, ReportType -from generalresearch.models.custom_types import ( - AwareDatetimeISO, - UUIDStr, - UUIDStrCoerce, -) -from generalresearch.pg_helper import PostgresConfig -from generalresearch.redis_helper import RedisConfig from generalresearch.utils.enum import ReprEnumMeta if TYPE_CHECKING: from generalresearch.incite.base import GRLDatasets + from generalresearch.incite.mergers.foundations.enriched_session import ( + EnrichedSessionMerge, + ) + from generalresearch.incite.mergers.foundations.enriched_wall import ( + EnrichedWallMerge, + ) + from generalresearch.managers.gr.authentication import ( + GRUserManager, + ) + from generalresearch.managers.gr.business import BusinessManager + from generalresearch.managers.gr.team import MembershipManager + from generalresearch.managers.thl.product import ProductManager + from generalresearch.models.custom_types import ( + AwareDatetimeISO, + UUIDStr, + UUIDStrCoerce, + ) from generalresearch.models.gr.authentication import GRUser from generalresearch.models.gr.business import Business from generalresearch.models.thl.product import Product + from generalresearch.pg_helper import PostgresConfig + from generalresearch.redis_helper import RedisConfig class MembershipPrivilege(Enum, metaclass=ReprEnumMeta): @@ -120,36 +126,17 @@ class Team(BaseModel): # --- Prefetch Methods --- - def prefetch_memberships(self, pg_config: PostgresConfig) -> None: - from generalresearch.managers.gr.team import MembershipManager - - mm = MembershipManager(pg_config=pg_config) - self.memberships = mm.get_by_team_id(team_id=self.id) - - def prefetch_gr_users( - self, pg_config: PostgresConfig, redis_config: RedisConfig - ) -> None: - from generalresearch.managers.gr.authentication import ( - GRUserManager, - ) - - gr_um = GRUserManager(pg_config=pg_config, redis_config=redis_config) - - self.gr_users = gr_um.get_by_team(team_id=self.id) - - def prefetch_businesses( - self, pg_config: PostgresConfig, redis_config: RedisConfig - ) -> None: - from generalresearch.managers.gr.business import BusinessManager + def prefetch_memberships(self, membership_manager: MembershipManager) -> None: + self.memberships = membership_manager.get_by_team_id(team_id=self.id) - bm = BusinessManager(pg_config=pg_config, redis_config=redis_config) - self.businesses = bm.get_by_team(team_id=self.id) + def prefetch_gr_users(self, gr_user_manager: GRUserManager) -> None: + self.gr_users = gr_user_manager.get_by_team(team_id=self.id) - def prefetch_products(self, thl_pg_config: PostgresConfig) -> None: - from generalresearch.managers.thl.product import ProductManager + def prefetch_businesses(self, business_manager: BusinessManager) -> None: + self.businesses = business_manager.get_by_team(team_id=self.id) - pm = ProductManager(pg_config=thl_pg_config) - self.products = pm.fetch_uuids(team_uuids=[self.uuid]) + def prefetch_products(self, product_manager: ProductManager) -> None: + self.products = product_manager.fetch_uuids(team_uuids=[self.uuid]) # --- Prebuild Methods --- diff --git a/generalresearch/models/innovate/question.py b/generalresearch/models/innovate/question.py index f5a4846..fc89524 100644 --- a/generalresearch/models/innovate/question.py +++ b/generalresearch/models/innovate/question.py @@ -9,13 +9,13 @@ 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.innovate import InnovateQuestionID from generalresearch.models.thl.profiling.marketplace import ( MarketplaceQuestion, MarketplaceUserQuestionAnswer, ) if TYPE_CHECKING: + from generalresearch.models.innovate import InnovateQuestionID from generalresearch.models.thl.profiling.upk_question import ( UpkQuestion, ) diff --git a/generalresearch/models/innovate/survey.py b/generalresearch/models/innovate/survey.py index d07f960..60921df 100644 --- a/generalresearch/models/innovate/survey.py +++ b/generalresearch/models/innovate/survey.py @@ -6,6 +6,7 @@ from datetime import UTC, date from decimal import Decimal from functools import cached_property from typing import ( + TYPE_CHECKING, Annotated, Any, Literal, @@ -26,20 +27,12 @@ from generalresearch.locales import Localelator from generalresearch.models import ( LogicalOperator, Source, - TaskCalculationType, -) -from generalresearch.models.custom_types import ( - AlphaNumStrSet, - AwareDatetimeISO, - CoercedStr, - DeviceTypes, ) from generalresearch.models.innovate import ( InnovateDuplicateCheckLevel, InnovateQuotaStatus, InnovateStatus, ) -from generalresearch.models.innovate.question import InnovateQuestionID from generalresearch.models.thl.demographics import Gender from generalresearch.models.thl.survey import MarketplaceTask from generalresearch.models.thl.survey.condition import ( @@ -47,6 +40,18 @@ from generalresearch.models.thl.survey.condition import ( MarketplaceCondition, ) +if TYPE_CHECKING: + from generalresearch.models import ( + TaskCalculationType, + ) + from generalresearch.models.custom_types import ( + AlphaNumStrSet, + AwareDatetimeISO, + CoercedStr, + DeviceTypes, + ) + from generalresearch.models.innovate.question import InnovateQuestionID + logging.basicConfig() logger = logging.getLogger() logger.setLevel(logging.INFO) diff --git a/generalresearch/models/innovate/task_collection.py b/generalresearch/models/innovate/task_collection.py index 7bf9d0f..a647d30 100644 --- a/generalresearch/models/innovate/task_collection.py +++ b/generalresearch/models/innovate/task_collection.py @@ -1,16 +1,20 @@ from __future__ import annotations +from typing import TYPE_CHECKING + import pandas as pd from pandera.pandas import Check, Column, DataFrameSchema, Index from generalresearch.locales import Localelator from generalresearch.models.innovate import InnovateStatus -from generalresearch.models.innovate.survey import InnovateSurvey from generalresearch.models.thl.survey.task_collection import ( TaskCollection, create_empty_df_from_schema, ) +if TYPE_CHECKING: + from generalresearch.models.innovate.survey import InnovateSurvey + COUNTRY_ISOS: set[str] = Localelator().get_all_countries() LANGUAGE_ISOS: set[str] = Localelator().get_all_languages() diff --git a/generalresearch/models/legacy/bucket.py b/generalresearch/models/legacy/bucket.py index 2650b90..f20a769 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, Self +from typing import TYPE_CHECKING, Any, Literal, Self from pydantic import ( BaseModel, @@ -16,13 +16,15 @@ from pydantic import ( ) from generalresearch.models import Source -from generalresearch.models.custom_types import ( - HttpsUrl, - PropertyCode, - UUIDStr, -) from generalresearch.models.thl.stats import StatisticalSummary +if TYPE_CHECKING: + from generalresearch.models.custom_types import ( + HttpsUrl, + PropertyCode, + UUIDStr, + ) + logger = logging.getLogger() Eligibility = Literal["conditional", "unconditional", "ineligible"] diff --git a/generalresearch/models/legacy/offerwall.py b/generalresearch/models/legacy/offerwall.py index da28663..c150dcb 100644 --- a/generalresearch/models/legacy/offerwall.py +++ b/generalresearch/models/legacy/offerwall.py @@ -1,27 +1,33 @@ from __future__ import annotations +from typing import TYPE_CHECKING + from pydantic import BaseModel, ConfigDict, Field, NonNegativeInt -from generalresearch.models.custom_types import UUIDStr -from generalresearch.models.legacy.bucket import ( - BucketBase, - MarketplaceBucket, - OneShotOfferwallBucket, - OneShotSoftPairOfferwallBucket, - SingleEntryBucket, - SoftPairBucket, - TimeBucksBucket, - TopNBucket, - TopNPlusBucket, - TopNPlusRecontactBucket, - WXETOfferwallBucket, -) from generalresearch.models.legacy.definitions import OfferwallReason from generalresearch.models.thl.payout_format import ( PayoutFormatField, - PayoutFormatType, ) -from generalresearch.models.thl.profiling.upk_question import UpkQuestion + +if TYPE_CHECKING: + from generalresearch.models.custom_types import UUIDStr + from generalresearch.models.legacy.bucket import ( + BucketBase, + MarketplaceBucket, + OneShotOfferwallBucket, + OneShotSoftPairOfferwallBucket, + SingleEntryBucket, + SoftPairBucket, + TimeBucksBucket, + TopNBucket, + TopNPlusBucket, + TopNPlusRecontactBucket, + WXETOfferwallBucket, + ) + from generalresearch.models.thl.payout_format import ( + PayoutFormatType, + ) + from generalresearch.models.thl.profiling.upk_question import UpkQuestion """ Not Done: diff --git a/generalresearch/models/legacy/questions.py b/generalresearch/models/legacy/questions.py index 9f37837..bebd28f 100644 --- a/generalresearch/models/legacy/questions.py +++ b/generalresearch/models/legacy/questions.py @@ -15,19 +15,19 @@ from pydantic import ( ) from sentry_sdk import capture_exception -from generalresearch.models.custom_types import UUIDStr from generalresearch.models.legacy.api_status import StatusResponse -from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestionOut, -) -from generalresearch.models.thl.session import Wall -from generalresearch.models.thl.user import User if TYPE_CHECKING: from generalresearch.managers.thl.user_manager.user_manager import ( UserManager, ) from generalresearch.managers.thl.wall import WallManager + from generalresearch.models.custom_types import UUIDStr + from generalresearch.models.thl.profiling.upk_question import ( + UpkQuestionOut, + ) + from generalresearch.models.thl.session import Wall + from generalresearch.models.thl.user import User class UpkQuestionResponse(StatusResponse): diff --git a/generalresearch/models/lucid/question.py b/generalresearch/models/lucid/question.py index af3420e..98f535b 100644 --- a/generalresearch/models/lucid/question.py +++ b/generalresearch/models/lucid/question.py @@ -7,12 +7,12 @@ 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.lucid import LucidQuestionIdType from generalresearch.models.thl.profiling.marketplace import ( MarketplaceQuestion, ) if TYPE_CHECKING: + from generalresearch.models.lucid import LucidQuestionIdType from generalresearch.models.thl.profiling.upk_question import ( UpkQuestion, ) diff --git a/generalresearch/models/lucid/survey.py b/generalresearch/models/lucid/survey.py index 4b1bb98..0f03e31 100644 --- a/generalresearch/models/lucid/survey.py +++ b/generalresearch/models/lucid/survey.py @@ -1,22 +1,24 @@ from __future__ import annotations -from typing import Any, Self +from typing import TYPE_CHECKING, Any, Self from pydantic import BaseModel, ConfigDict, Field, NonNegativeInt from generalresearch.models import Source -from generalresearch.models.custom_types import ( - AwareDatetimeISO, - BigAutoInteger, - CoercedStr, - UUIDStr, -) -from generalresearch.models.thl.locales import CountryISO, LanguageISO from generalresearch.models.thl.survey.condition import ( ConditionValueType, MarketplaceCondition, ) +if TYPE_CHECKING: + from generalresearch.models.custom_types import ( + AwareDatetimeISO, + BigAutoInteger, + CoercedStr, + UUIDStr, + ) + from generalresearch.models.thl.locales import CountryISO, LanguageISO + class LucidCondition(MarketplaceCondition): model_config = ConfigDict(populate_by_name=True, frozen=False, extra="ignore") diff --git a/generalresearch/models/morning/question.py b/generalresearch/models/morning/question.py index b64a44a..748fcc6 100644 --- a/generalresearch/models/morning/question.py +++ b/generalresearch/models/morning/question.py @@ -1,18 +1,20 @@ import json from enum import StrEnum -from typing import Any, Literal, Self +from typing import TYPE_CHECKING, Any, Literal, Self 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.morning import MorningQuestionID from generalresearch.models.thl.profiling.marketplace import ( MarketplaceQuestion, MarketplaceUserQuestionAnswer, ) +if TYPE_CHECKING: + from generalresearch.models.morning import MorningQuestionID + # todo: we could validate that the country_iso / language_iso exists ... locale_helper = Localelator() diff --git a/generalresearch/models/morning/survey.py b/generalresearch/models/morning/survey.py index 91d1bce..1e217f6 100644 --- a/generalresearch/models/morning/survey.py +++ b/generalresearch/models/morning/survey.py @@ -6,6 +6,7 @@ from datetime import UTC from decimal import Decimal from functools import cached_property from typing import ( + TYPE_CHECKING, Annotated, Any, Literal, @@ -25,24 +26,27 @@ from pydantic import ( from generalresearch.locales import Localelator from generalresearch.models import Source -from generalresearch.models.custom_types import ( - AwareDatetimeISO, - UUIDStrCoerce, -) -from generalresearch.models.morning import MorningQuestionID, MorningStatus -from generalresearch.models.morning.question import MorningQuestion +from generalresearch.models.morning import MorningStatus from generalresearch.models.thl.demographics import Gender -from generalresearch.models.thl.locales import ( - CountryISO, - CountryISOs, - LanguageISOs, -) from generalresearch.models.thl.survey import MarketplaceTask from generalresearch.models.thl.survey.condition import ( ConditionValueType, MarketplaceCondition, ) +if TYPE_CHECKING: + from generalresearch.models.custom_types import ( + AwareDatetimeISO, + UUIDStrCoerce, + ) + from generalresearch.models.morning import MorningQuestionID + from generalresearch.models.morning.question import MorningQuestion + from generalresearch.models.thl.locales import ( + CountryISO, + CountryISOs, + LanguageISOs, + ) + logging.basicConfig() logger = logging.getLogger() logger.setLevel(logging.INFO) diff --git a/generalresearch/models/morning/task_collection.py b/generalresearch/models/morning/task_collection.py index 9303a2f..1117937 100644 --- a/generalresearch/models/morning/task_collection.py +++ b/generalresearch/models/morning/task_collection.py @@ -1,16 +1,20 @@ from __future__ import annotations +from typing import TYPE_CHECKING + import pandas as pd from pandera.pandas import Check, Column, DataFrameSchema, Index from generalresearch.locales import Localelator from generalresearch.models.morning import MorningStatus -from generalresearch.models.morning.survey import MorningBid from generalresearch.models.thl.survey.task_collection import ( TaskCollection, create_empty_df_from_schema, ) +if TYPE_CHECKING: + from generalresearch.models.morning.survey import MorningBid + COUNTRY_ISOS: set[str] = Localelator().get_all_countries() LANGUAGE_ISOS: set[str] = Localelator().get_all_languages() diff --git a/generalresearch/models/network/label.py b/generalresearch/models/network/label.py index b8fe4b0..c27f36f 100644 --- a/generalresearch/models/network/label.py +++ b/generalresearch/models/network/label.py @@ -3,6 +3,7 @@ from __future__ import annotations import ipaddress from enum import StrEnum from ipaddress import IPv4Network, IPv6Network +from typing import TYPE_CHECKING from pydantic import ( BaseModel, @@ -13,10 +14,10 @@ from pydantic import ( field_validator, ) -from generalresearch.models.custom_types import ( - AwareDatetimeISO, - now_utc_factory, -) +from generalresearch.models.custom_types import now_utc_factory + +if TYPE_CHECKING: + from generalresearch.models.custom_types import AwareDatetimeISO class IPTrustClass(StrEnum): diff --git a/generalresearch/models/network/mtr/command.py b/generalresearch/models/network/mtr/command.py index 220bfc7..7e74f20 100644 --- a/generalresearch/models/network/mtr/command.py +++ b/generalresearch/models/network/mtr/command.py @@ -5,9 +5,9 @@ from typing import TYPE_CHECKING from generalresearch.models.network.definitions import IPProtocol from generalresearch.models.network.mtr.parser import parse_mtr_output -from generalresearch.models.network.mtr.result import MTRResult if TYPE_CHECKING: + from generalresearch.models.network.mtr.result import MTRResult from generalresearch.models.network.tool_run_command import MTRRunCommand SUPPORTED_PROTOCOLS = { diff --git a/generalresearch/models/network/mtr/execute.py b/generalresearch/models/network/mtr/execute.py index c5b3c5c..1a7c963 100644 --- a/generalresearch/models/network/mtr/execute.py +++ b/generalresearch/models/network/mtr/execute.py @@ -1,21 +1,29 @@ from __future__ import annotations from datetime import UTC, datetime +from typing import TYPE_CHECKING from uuid import uuid4 -from generalresearch.models.custom_types import UUIDStr from generalresearch.models.network.definitions import IPProtocol from generalresearch.models.network.mtr.command import ( get_mtr_version, run_mtr, ) -from generalresearch.models.network.tool_run import MTRRun, Status, ToolClass, ToolName +from generalresearch.models.network.tool_run import ( + MTRRun, + Status, + ToolClass, + ToolName, +) from generalresearch.models.network.tool_run_command import ( MTRRunCommand, MTRRunCommandOptions, ) from generalresearch.models.network.utils import get_source_ip +if TYPE_CHECKING: + from generalresearch.models.custom_types import UUIDStr + def execute_mtr( ip: str, diff --git a/generalresearch/models/network/mtr/parser.py b/generalresearch/models/network/mtr/parser.py index c29439e..a0f8998 100644 --- a/generalresearch/models/network/mtr/parser.py +++ b/generalresearch/models/network/mtr/parser.py @@ -1,8 +1,11 @@ import json +from typing import TYPE_CHECKING -from generalresearch.models.network.definitions import IPProtocol from generalresearch.models.network.mtr.result import MTRResult +if TYPE_CHECKING: + from generalresearch.models.network.definitions import IPProtocol + def parse_mtr_output(raw: str, port: int, protocol: IPProtocol) -> MTRResult: data = parse_mtr_raw_output(raw) diff --git a/generalresearch/models/network/mtr/result.py b/generalresearch/models/network/mtr/result.py index 34de845..d17136c 100644 --- a/generalresearch/models/network/mtr/result.py +++ b/generalresearch/models/network/mtr/result.py @@ -3,6 +3,7 @@ from __future__ import annotations import re from functools import cached_property from ipaddress import ip_address +from typing import TYPE_CHECKING import tldextract from pydantic import ( @@ -14,7 +15,13 @@ from pydantic import ( model_validator, ) -from generalresearch.models.network.definitions import IPKind, IPProtocol, get_ip_kind +from generalresearch.models.network.definitions import ( + IPProtocol, + get_ip_kind, +) + +if TYPE_CHECKING: + from generalresearch.models.network.definitions import IPKind HOST_RE = re.compile(r"^(?P.+?) \((?P[^)]+)\)$") diff --git a/generalresearch/models/network/nmap/command.py b/generalresearch/models/network/nmap/command.py index 42b7178..3509b8d 100644 --- a/generalresearch/models/network/nmap/command.py +++ b/generalresearch/models/network/nmap/command.py @@ -4,9 +4,9 @@ import subprocess from typing import TYPE_CHECKING from generalresearch.models.network.nmap.parser import parse_nmap_xml -from generalresearch.models.network.nmap.result import NmapResult if TYPE_CHECKING: + from generalresearch.models.network.nmap.result import NmapResult from generalresearch.models.network.tool_run_command import NmapRunCommand diff --git a/generalresearch/models/network/nmap/execute.py b/generalresearch/models/network/nmap/execute.py index 6c05c89..e3610d9 100644 --- a/generalresearch/models/network/nmap/execute.py +++ b/generalresearch/models/network/nmap/execute.py @@ -1,15 +1,23 @@ from __future__ import annotations +from typing import TYPE_CHECKING from uuid import uuid4 -from generalresearch.models.custom_types import UUIDStr from generalresearch.models.network.nmap.command import run_nmap -from generalresearch.models.network.tool_run import NmapRun, Status, ToolClass, ToolName +from generalresearch.models.network.tool_run import ( + NmapRun, + Status, + ToolClass, + ToolName, +) from generalresearch.models.network.tool_run_command import ( NmapRunCommand, NmapRunCommandOptions, ) +if TYPE_CHECKING: + from generalresearch.models.custom_types import UUIDStr + def execute_nmap( ip: str, diff --git a/generalresearch/models/network/nmap/result.py b/generalresearch/models/network/nmap/result.py index 4552e15..57c2e8b 100644 --- a/generalresearch/models/network/nmap/result.py +++ b/generalresearch/models/network/nmap/result.py @@ -4,13 +4,15 @@ import json from datetime import timedelta from enum import StrEnum from functools import cached_property -from typing import Any, Literal +from typing import TYPE_CHECKING, Any, Literal from pydantic import BaseModel, Field, computed_field -from generalresearch.models.custom_types import AwareDatetimeISO, IPvAnyAddressStr from generalresearch.models.network.definitions import IPProtocol +if TYPE_CHECKING: + from generalresearch.models.custom_types import AwareDatetimeISO, IPvAnyAddressStr + class PortState(StrEnum): OPEN = "open" diff --git a/generalresearch/models/network/rdns/command.py b/generalresearch/models/network/rdns/command.py index c63d6d2..2250449 100644 --- a/generalresearch/models/network/rdns/command.py +++ b/generalresearch/models/network/rdns/command.py @@ -2,9 +2,9 @@ import subprocess from typing import TYPE_CHECKING from generalresearch.models.network.rdns.parser import parse_rdns_output -from generalresearch.models.network.rdns.result import RDNSResult if TYPE_CHECKING: + from generalresearch.models.network.rdns.result import RDNSResult from generalresearch.models.network.tool_run_command import RDNSRunCommand diff --git a/generalresearch/models/network/rdns/execute.py b/generalresearch/models/network/rdns/execute.py index d6de84b..6c14f77 100644 --- a/generalresearch/models/network/rdns/execute.py +++ b/generalresearch/models/network/rdns/execute.py @@ -1,9 +1,9 @@ from __future__ import annotations from datetime import UTC, datetime +from typing import TYPE_CHECKING from uuid import uuid4 -from generalresearch.models.custom_types import UUIDStr from generalresearch.models.network.rdns.command import ( get_dig_version, run_rdns, @@ -19,6 +19,9 @@ from generalresearch.models.network.tool_run_command import ( RDNSRunCommandOptions, ) +if TYPE_CHECKING: + from generalresearch.models.custom_types import UUIDStr + def execute_rdns(ip: str, scan_group_id: UUIDStr | None = None): started_at = datetime.now(tz=UTC) diff --git a/generalresearch/models/network/rdns/parser.py b/generalresearch/models/network/rdns/parser.py index 31a5ed6..dc33997 100644 --- a/generalresearch/models/network/rdns/parser.py +++ b/generalresearch/models/network/rdns/parser.py @@ -1,9 +1,12 @@ import ipaddress import re +from typing import TYPE_CHECKING -from generalresearch.models.custom_types import IPvAnyAddressStr from generalresearch.models.network.rdns.result import RDNSResult +if TYPE_CHECKING: + from generalresearch.models.custom_types import IPvAnyAddressStr + PTR_RE = re.compile(r"\sPTR\s+([^\s]+)\.") diff --git a/generalresearch/models/network/rdns/result.py b/generalresearch/models/network/rdns/result.py index 46af643..6845775 100644 --- a/generalresearch/models/network/rdns/result.py +++ b/generalresearch/models/network/rdns/result.py @@ -2,11 +2,13 @@ from __future__ import annotations import json from functools import cached_property +from typing import TYPE_CHECKING import tldextract from pydantic import BaseModel, Field, computed_field, model_validator -from generalresearch.models.custom_types import IPvAnyAddressStr +if TYPE_CHECKING: + from generalresearch.models.custom_types import IPvAnyAddressStr class RDNSResult(BaseModel): diff --git a/generalresearch/models/network/tool_run.py b/generalresearch/models/network/tool_run.py index c49ffc0..8479f15 100644 --- a/generalresearch/models/network/tool_run.py +++ b/generalresearch/models/network/tool_run.py @@ -1,25 +1,26 @@ from __future__ import annotations from enum import StrEnum -from typing import Literal +from typing import TYPE_CHECKING, Literal from uuid import uuid4 from pydantic import BaseModel, Field, PositiveInt -from generalresearch.models.custom_types import ( - AwareDatetimeISO, - IPvAnyAddressStr, - UUIDStr, -) -from generalresearch.models.network.mtr.result import MTRResult -from generalresearch.models.network.nmap.result import NmapResult -from generalresearch.models.network.rdns.result import RDNSResult -from generalresearch.models.network.tool_run_command import ( - MTRRunCommand, - NmapRunCommand, - RDNSRunCommand, - ToolRunCommand, -) +if TYPE_CHECKING: + from generalresearch.models.custom_types import ( + AwareDatetimeISO, + IPvAnyAddressStr, + UUIDStr, + ) + from generalresearch.models.network.mtr.result import MTRResult + from generalresearch.models.network.nmap.result import NmapResult + from generalresearch.models.network.rdns.result import RDNSResult + from generalresearch.models.network.tool_run_command import ( + MTRRunCommand, + NmapRunCommand, + RDNSRunCommand, + ToolRunCommand, + ) class ToolClass(StrEnum): diff --git a/generalresearch/models/network/tool_run_command.py b/generalresearch/models/network/tool_run_command.py index 6f22d6b..b07b811 100644 --- a/generalresearch/models/network/tool_run_command.py +++ b/generalresearch/models/network/tool_run_command.py @@ -1,12 +1,14 @@ from __future__ import annotations -from typing import Literal +from typing import TYPE_CHECKING, Literal from pydantic import BaseModel, Field -from generalresearch.models.custom_types import IPvAnyAddressStr from generalresearch.models.network.definitions import IPProtocol +if TYPE_CHECKING: + from generalresearch.models.custom_types import IPvAnyAddressStr + class ToolRunCommand(BaseModel): command: str = Field() diff --git a/generalresearch/models/precision/question.py b/generalresearch/models/precision/question.py index f532998..ba17361 100644 --- a/generalresearch/models/precision/question.py +++ b/generalresearch/models/precision/question.py @@ -9,13 +9,13 @@ 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.precision import PrecisionQuestionID from generalresearch.models.thl.profiling.marketplace import ( MarketplaceQuestion, MarketplaceUserQuestionAnswer, ) if TYPE_CHECKING: + from generalresearch.models.precision import PrecisionQuestionID from generalresearch.models.thl.profiling.upk_question import ( UpkQuestion, ) diff --git a/generalresearch/models/precision/survey.py b/generalresearch/models/precision/survey.py index bf9e83e..fa30882 100644 --- a/generalresearch/models/precision/survey.py +++ b/generalresearch/models/precision/survey.py @@ -3,7 +3,7 @@ from __future__ import annotations import json from datetime import UTC from functools import cached_property -from typing import Annotated, Any, Literal, Self +from typing import TYPE_CHECKING, Annotated, Any, Literal, Self from more_itertools import flatten from pydantic import ( @@ -16,14 +16,7 @@ from pydantic import ( ) from generalresearch.models import Source -from generalresearch.models.custom_types import ( - AlphaNumStrSet, - AwareDatetimeISO, - CoercedStr, - DeviceTypes, - UUIDStrCoerce, -) -from generalresearch.models.precision import PrecisionQuestionID, PrecisionStatus +from generalresearch.models.precision import PrecisionStatus from generalresearch.models.thl.demographics import Gender from generalresearch.models.thl.survey import MarketplaceTask from generalresearch.models.thl.survey.condition import ( @@ -31,6 +24,16 @@ from generalresearch.models.thl.survey.condition import ( MarketplaceCondition, ) +if TYPE_CHECKING: + from generalresearch.models.custom_types import ( + AlphaNumStrSet, + AwareDatetimeISO, + CoercedStr, + DeviceTypes, + UUIDStrCoerce, + ) + from generalresearch.models.precision import PrecisionQuestionID + class PrecisionCondition(MarketplaceCondition): question_id: PrecisionQuestionID | None = Field() diff --git a/generalresearch/models/precision/task_collection.py b/generalresearch/models/precision/task_collection.py index c8db2af..71241ad 100644 --- a/generalresearch/models/precision/task_collection.py +++ b/generalresearch/models/precision/task_collection.py @@ -1,16 +1,18 @@ -from typing import Any +from typing import TYPE_CHECKING, Any import pandas as pd from pandera.pandas import Check, Column, DataFrameSchema, Index from generalresearch.locales import Localelator from generalresearch.models.precision import PrecisionStatus -from generalresearch.models.precision.survey import PrecisionSurvey from generalresearch.models.thl.survey.task_collection import ( TaskCollection, create_empty_df_from_schema, ) +if TYPE_CHECKING: + from generalresearch.models.precision.survey import PrecisionSurvey + COUNTRY_ISOS = Localelator().get_all_countries() LANGUAGE_ISOS = Localelator().get_all_languages() diff --git a/generalresearch/models/prodege/question.py b/generalresearch/models/prodege/question.py index 58aed67..c43b51a 100644 --- a/generalresearch/models/prodege/question.py +++ b/generalresearch/models/prodege/question.py @@ -19,11 +19,11 @@ from pydantic import ( from generalresearch.locales import Localelator from generalresearch.models import MAX_INT32, Source -from generalresearch.models.custom_types import AwareDatetimeISO -from generalresearch.models.prodege import ProdegeQuestionIdType from generalresearch.models.thl.profiling.marketplace import MarketplaceQuestion if TYPE_CHECKING: + from generalresearch.models.custom_types import AwareDatetimeISO + from generalresearch.models.prodege import ProdegeQuestionIdType from generalresearch.models.thl.profiling.upk_question import ( UpkQuestion, ) diff --git a/generalresearch/models/prodege/survey.py b/generalresearch/models/prodege/survey.py index 5d0369a..7e56a9c 100644 --- a/generalresearch/models/prodege/survey.py +++ b/generalresearch/models/prodege/survey.py @@ -7,7 +7,7 @@ from collections import defaultdict from datetime import UTC, datetime from decimal import Decimal from functools import cached_property -from typing import Any, Literal +from typing import TYPE_CHECKING, Any, Literal from pydantic import ( BaseModel, @@ -21,18 +21,9 @@ from pydantic import ( from generalresearch.locales import Localelator from generalresearch.models import LogicalOperator, Source, TaskCalculationType -from generalresearch.models.custom_types import ( - AlphaNumStrSet, - AwareDatetimeISO, - CoercedStr, - InclExcl, - UUIDStr, -) from generalresearch.models.prodege import ( ProdegePastParticipationType, - ProdegeQuestionIdType, ProdegeStatus, - ProdgeRedirectStatus, ) from generalresearch.models.prodege.definitions import PG_COUNTRY_TO_ISO from generalresearch.models.thl.demographics import Gender @@ -42,6 +33,19 @@ from generalresearch.models.thl.survey.condition import ( MarketplaceCondition, ) +if TYPE_CHECKING: + from generalresearch.models.custom_types import ( + AlphaNumStrSet, + AwareDatetimeISO, + CoercedStr, + InclExcl, + UUIDStr, + ) + from generalresearch.models.prodege import ( + ProdegeQuestionIdType, + ProdgeRedirectStatus, + ) + logging.basicConfig() logger = logging.getLogger() logger.setLevel(logging.INFO) diff --git a/generalresearch/models/prodege/task_collection.py b/generalresearch/models/prodege/task_collection.py index 9f6a81b..774fc7b 100644 --- a/generalresearch/models/prodege/task_collection.py +++ b/generalresearch/models/prodege/task_collection.py @@ -1,16 +1,18 @@ -from typing import Any +from typing import TYPE_CHECKING, Any import pandas as pd from pandera.pandas import Check, Column, DataFrameSchema, Index from generalresearch.locales import Localelator from generalresearch.models.prodege import ProdegeStatus -from generalresearch.models.prodege.survey import ProdegeSurvey from generalresearch.models.thl.survey.task_collection import ( TaskCollection, create_empty_df_from_schema, ) +if TYPE_CHECKING: + from generalresearch.models.prodege.survey import ProdegeSurvey + COUNTRY_ISOS = Localelator().get_all_countries() LANGUAGE_ISOS = Localelator().get_all_languages() diff --git a/generalresearch/models/repdata/question.py b/generalresearch/models/repdata/question.py index 4fa2d22..8cb1fa7 100644 --- a/generalresearch/models/repdata/question.py +++ b/generalresearch/models/repdata/question.py @@ -18,10 +18,10 @@ from pydantic import ( ) from generalresearch.models import MAX_INT32, Source -from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.thl.profiling.marketplace import MarketplaceQuestion if TYPE_CHECKING: + from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.thl.profiling.upk_question import ( UpkQuestion, ) diff --git a/generalresearch/models/repdata/survey.py b/generalresearch/models/repdata/survey.py index c5b0730..cea61ed 100644 --- a/generalresearch/models/repdata/survey.py +++ b/generalresearch/models/repdata/survey.py @@ -6,7 +6,7 @@ import logging from datetime import UTC, datetime from decimal import Decimal from functools import cached_property -from typing import Any, Literal, Self +from typing import TYPE_CHECKING, Any, Literal, Self from uuid import UUID from pydantic import ( @@ -27,11 +27,6 @@ from generalresearch.models import ( Source, TaskCalculationType, ) -from generalresearch.models.custom_types import ( - AwareDatetimeISO, - CoercedStr, - UUIDStr, -) from generalresearch.models.repdata import RepDataStatus from generalresearch.models.thl.demographics import Gender from generalresearch.models.thl.survey import MarketplaceTask @@ -40,6 +35,13 @@ from generalresearch.models.thl.survey.condition import ( MarketplaceCondition, ) +if TYPE_CHECKING: + from generalresearch.models.custom_types import ( + AwareDatetimeISO, + CoercedStr, + UUIDStr, + ) + logging.basicConfig() logger = logging.getLogger() logger.setLevel(logging.INFO) diff --git a/generalresearch/models/repdata/task_collection.py b/generalresearch/models/repdata/task_collection.py index 5b9a4ba..04d79bd 100644 --- a/generalresearch/models/repdata/task_collection.py +++ b/generalresearch/models/repdata/task_collection.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import Any +from typing import TYPE_CHECKING, Any import pandas as pd from pandera.pandas import Check, Column, DataFrameSchema, Index @@ -8,12 +8,14 @@ from pandera.pandas import Check, Column, DataFrameSchema, Index from generalresearch.locales import Localelator from generalresearch.models import TaskCalculationType from generalresearch.models.repdata import RepDataStatus -from generalresearch.models.repdata.survey import RepDataSurveyHashed from generalresearch.models.thl.survey.task_collection import ( TaskCollection, create_empty_df_from_schema, ) +if TYPE_CHECKING: + from generalresearch.models.repdata.survey import RepDataSurveyHashed + COUNTRY_ISOS = Localelator().get_all_countries() LANGUAGE_ISOS = Localelator().get_all_languages() diff --git a/generalresearch/models/sago/question.py b/generalresearch/models/sago/question.py index 291214f..cf9ea19 100644 --- a/generalresearch/models/sago/question.py +++ b/generalresearch/models/sago/question.py @@ -19,10 +19,10 @@ from pydantic import ( ) from generalresearch.models import MAX_INT32, Source, string_utils -from generalresearch.models.custom_types import AwareDatetimeISO from generalresearch.models.thl.profiling.marketplace import MarketplaceQuestion if TYPE_CHECKING: + from generalresearch.models.custom_types import AwareDatetimeISO from generalresearch.models.thl.profiling.upk_question import ( UpkQuestion, ) diff --git a/generalresearch/models/sago/survey.py b/generalresearch/models/sago/survey.py index 8550cd3..c2f886a 100644 --- a/generalresearch/models/sago/survey.py +++ b/generalresearch/models/sago/survey.py @@ -5,7 +5,7 @@ import logging from datetime import UTC from decimal import Decimal from functools import cached_property -from typing import Annotated, Any, Literal, Self +from typing import TYPE_CHECKING, Annotated, Any, Literal, Self from more_itertools import flatten from pydantic import ( @@ -19,14 +19,6 @@ from pydantic import ( from generalresearch.locales import Localelator from generalresearch.models import LogicalOperator, Source -from generalresearch.models.custom_types import ( - AlphaNumStr, - AlphaNumStrSet, - AwareDatetimeISO, - CoercedStr, - DeviceTypes, - IPLikeStrSet, -) from generalresearch.models.sago import SagoStatus from generalresearch.models.thl.demographics import Gender from generalresearch.models.thl.survey import MarketplaceTask @@ -35,6 +27,16 @@ from generalresearch.models.thl.survey.condition import ( MarketplaceCondition, ) +if TYPE_CHECKING: + from generalresearch.models.custom_types import ( + AlphaNumStr, + AlphaNumStrSet, + AwareDatetimeISO, + CoercedStr, + DeviceTypes, + IPLikeStrSet, + ) + logging.basicConfig() logger = logging.getLogger() logger.setLevel(logging.INFO) diff --git a/generalresearch/models/sago/task_collection.py b/generalresearch/models/sago/task_collection.py index 2879d9c..490f7a0 100644 --- a/generalresearch/models/sago/task_collection.py +++ b/generalresearch/models/sago/task_collection.py @@ -1,18 +1,20 @@ from __future__ import annotations -from typing import Any +from typing import TYPE_CHECKING, Any import pandas as pd from pandera.pandas import Check, Column, DataFrameSchema, Index from generalresearch.locales import Localelator from generalresearch.models.sago import SagoStatus -from generalresearch.models.sago.survey import SagoSurvey from generalresearch.models.thl.survey.task_collection import ( TaskCollection, create_empty_df_from_schema, ) +if TYPE_CHECKING: + from generalresearch.models.sago.survey import SagoSurvey + COUNTRY_ISOS: set[str] = Localelator().get_all_countries() LANGUAGE_ISOS: set[str] = Localelator().get_all_languages() diff --git a/generalresearch/models/spectrum/question.py b/generalresearch/models/spectrum/question.py index 7add692..89fbeb3 100644 --- a/generalresearch/models/spectrum/question.py +++ b/generalresearch/models/spectrum/question.py @@ -19,13 +19,13 @@ from pydantic import ( ) from generalresearch.models import MAX_INT32, Source, string_utils -from generalresearch.models.custom_types import AwareDatetimeISO -from generalresearch.models.spectrum import SpectrumQuestionIdType from generalresearch.models.thl.profiling.marketplace import ( MarketplaceQuestion, ) if TYPE_CHECKING: + from generalresearch.models.custom_types import AwareDatetimeISO + from generalresearch.models.spectrum import SpectrumQuestionIdType from generalresearch.models.thl.profiling.upk_question import ( UpkQuestion, ) diff --git a/generalresearch/models/spectrum/survey.py b/generalresearch/models/spectrum/survey.py index 424d206..4daa00b 100644 --- a/generalresearch/models/spectrum/survey.py +++ b/generalresearch/models/spectrum/survey.py @@ -4,20 +4,13 @@ import json import logging from datetime import UTC from decimal import Decimal -from typing import Any, Literal, Self +from typing import TYPE_CHECKING, Any, Literal, Self 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.custom_types import ( - AlphaNumStr, - AlphaNumStrSet, - AwareDatetimeISO, - CoercedStr, - UUIDStrSet, -) from generalresearch.models.spectrum import SpectrumStatus from generalresearch.models.thl.demographics import Gender from generalresearch.models.thl.survey import MarketplaceTask @@ -26,6 +19,15 @@ from generalresearch.models.thl.survey.condition import ( MarketplaceCondition, ) +if TYPE_CHECKING: + from generalresearch.models.custom_types import ( + AlphaNumStr, + AlphaNumStrSet, + AwareDatetimeISO, + CoercedStr, + UUIDStrSet, + ) + logging.basicConfig() logger = logging.getLogger() logger.setLevel(logging.INFO) diff --git a/generalresearch/models/spectrum/task_collection.py b/generalresearch/models/spectrum/task_collection.py index 8ca5a93..d909292 100644 --- a/generalresearch/models/spectrum/task_collection.py +++ b/generalresearch/models/spectrum/task_collection.py @@ -1,17 +1,21 @@ from __future__ import annotations +from typing import TYPE_CHECKING + 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.spectrum import SpectrumStatus -from generalresearch.models.spectrum.survey import SpectrumSurvey from generalresearch.models.thl.survey.task_collection import ( TaskCollection, create_empty_df_from_schema, ) +if TYPE_CHECKING: + from generalresearch.models.spectrum.survey import SpectrumSurvey + COUNTRY_ISOS: set[str] = Localelator().get_all_countries() LANGUAGE_ISOS: set[str] = Localelator().get_all_languages() diff --git a/generalresearch/models/thl/category.py b/generalresearch/models/thl/category.py index ebfc840..32841a5 100644 --- a/generalresearch/models/thl/category.py +++ b/generalresearch/models/thl/category.py @@ -1,11 +1,12 @@ from __future__ import annotations -from typing import Any, Self +from typing import TYPE_CHECKING, Any, Self from uuid import uuid4 from pydantic import BaseModel, Field, PositiveInt, model_validator -from generalresearch.models.custom_types import UUIDStr +if TYPE_CHECKING: + from generalresearch.models.custom_types import UUIDStr class Category(BaseModel, frozen=True): diff --git a/generalresearch/models/thl/contest/__init__.py b/generalresearch/models/thl/contest/__init__.py index 65b28f4..f243ce4 100644 --- a/generalresearch/models/thl/contest/__init__.py +++ b/generalresearch/models/thl/contest/__init__.py @@ -1,7 +1,7 @@ from __future__ import annotations from datetime import UTC, datetime -from typing import Any, Self +from typing import TYPE_CHECKING, Any, Self from uuid import uuid4 from pydantic import ( @@ -12,10 +12,12 @@ from pydantic import ( model_validator, ) -from generalresearch.currency import USDCent -from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.thl.contest.definitions import ContestPrizeKind -from generalresearch.models.thl.user import User + +if TYPE_CHECKING: + from generalresearch.currency import USDCent + from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr + from generalresearch.models.thl.user import User class ContestEntryRule(BaseModel): diff --git a/generalresearch/models/thl/contest/contest.py b/generalresearch/models/thl/contest/contest.py index 6fc60f6..173d486 100644 --- a/generalresearch/models/thl/contest/contest.py +++ b/generalresearch/models/thl/contest/contest.py @@ -3,7 +3,7 @@ from __future__ import annotations import json from abc import ABC, abstractmethod from datetime import UTC, datetime -from typing import Any, Self +from typing import TYPE_CHECKING, Any, Self from uuid import uuid4 from pydantic import ( @@ -15,18 +15,22 @@ from pydantic import ( model_validator, ) -from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.thl.contest import ( ContestEndCondition, ContestPrize, - ContestWinner, ) from generalresearch.models.thl.contest.definitions import ( ContestEndReason, ContestStatus, ContestType, ) -from generalresearch.models.thl.locales import CountryISOs + +if TYPE_CHECKING: + from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr + from generalresearch.models.thl.contest import ( + ContestWinner, + ) + from generalresearch.models.thl.locales import CountryISOs class ContestBase(BaseModel, ABC): diff --git a/generalresearch/models/thl/contest/contest_entry.py b/generalresearch/models/thl/contest/contest_entry.py index b5f0ac3..a57b2df 100644 --- a/generalresearch/models/thl/contest/contest_entry.py +++ b/generalresearch/models/thl/contest/contest_entry.py @@ -1,7 +1,7 @@ from __future__ import annotations from datetime import UTC, datetime -from typing import Any +from typing import TYPE_CHECKING, Any from uuid import uuid4 from pydantic import ( @@ -12,9 +12,11 @@ from pydantic import ( ) from generalresearch.currency import USDCent -from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr -from generalresearch.models.thl.contest.definitions import ContestEntryType -from generalresearch.models.thl.user import User + +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 class ContestEntryCreate(BaseModel): diff --git a/generalresearch/models/thl/contest/leaderboard.py b/generalresearch/models/thl/contest/leaderboard.py index 080f356..efbbd0a 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 UTC, datetime, timedelta -from typing import Any, Literal, Self +from typing import TYPE_CHECKING, Any, Literal, Self from pydantic import ( ConfigDict, @@ -16,9 +16,6 @@ from generalresearch.currency import USDCent from generalresearch.decorators import LOG from generalresearch.managers.leaderboard import country_timezone from generalresearch.managers.leaderboard.manager import LeaderboardManager -from generalresearch.managers.thl.user_manager.user_manager import ( - UserManager, -) from generalresearch.models.thl.contest import ( ContestEndCondition, ContestPrize, @@ -42,6 +39,11 @@ from generalresearch.models.thl.leaderboard import ( LeaderboardFrequency, ) +if TYPE_CHECKING: + from generalresearch.managers.thl.user_manager.user_manager import ( + UserManager, + ) + class LeaderboardContestCreate(ContestBase): model_config = ConfigDict( diff --git a/generalresearch/models/thl/contest/milestone.py b/generalresearch/models/thl/contest/milestone.py index db5ba2f..e2ff2bc 100644 --- a/generalresearch/models/thl/contest/milestone.py +++ b/generalresearch/models/thl/contest/milestone.py @@ -2,7 +2,7 @@ from __future__ import annotations import logging from datetime import timedelta -from typing import Any, Literal, Self +from typing import TYPE_CHECKING, Any, Literal, Self from pydantic import ( BaseModel, @@ -13,7 +13,6 @@ from pydantic import ( ) from generalresearch.currency import USDCent -from generalresearch.models.custom_types import AwareDatetimeISO from generalresearch.models.thl.contest import ( ContestPrize, ) @@ -32,6 +31,9 @@ from generalresearch.models.thl.contest.definitions import ( ContestType, ) +if TYPE_CHECKING: + from generalresearch.models.custom_types import AwareDatetimeISO + logging.basicConfig() LOG = logging.getLogger() LOG.setLevel(logging.INFO) diff --git a/generalresearch/models/thl/contest/raffle.py b/generalresearch/models/thl/contest/raffle.py index 14b3fb6..b944740 100644 --- a/generalresearch/models/thl/contest/raffle.py +++ b/generalresearch/models/thl/contest/raffle.py @@ -4,7 +4,7 @@ import logging import random from collections import defaultdict from datetime import UTC, datetime -from typing import Any, Literal, Self +from typing import TYPE_CHECKING, Any, Literal, Self from pydantic import ( ConfigDict, @@ -20,7 +20,6 @@ from generalresearch.models.thl.contest import ( ContestEndCondition, ContestEntryRule, ContestPrize, - ContestWinner, ) from generalresearch.models.thl.contest.contest import ( Contest, @@ -28,7 +27,6 @@ from generalresearch.models.thl.contest.contest import ( ContestUserView, ) from generalresearch.models.thl.contest.contest_entry import ( - ContestEntry, ContestEntryType, ) from generalresearch.models.thl.contest.definitions import ( @@ -38,6 +36,14 @@ from generalresearch.models.thl.contest.definitions import ( ContestType, ) +if TYPE_CHECKING: + from generalresearch.models.thl.contest import ( + ContestWinner, + ) + from generalresearch.models.thl.contest.contest_entry import ( + ContestEntry, + ) + logging.basicConfig() LOG = logging.getLogger() LOG.setLevel(logging.INFO) diff --git a/generalresearch/models/thl/demographics.py b/generalresearch/models/thl/demographics.py index c11f8b2..ce4939c 100644 --- a/generalresearch/models/thl/demographics.py +++ b/generalresearch/models/thl/demographics.py @@ -8,9 +8,8 @@ from typing import TYPE_CHECKING, Any, Literal import numpy as np -from generalresearch.models.thl.locales import CountryISO - if TYPE_CHECKING: + from generalresearch.models.thl.locales import CountryISO from generalresearch.models.thl.survey import MarketplaceTask diff --git a/generalresearch/models/thl/finance.py b/generalresearch/models/thl/finance.py index 79a74a7..4b750da 100644 --- a/generalresearch/models/thl/finance.py +++ b/generalresearch/models/thl/finance.py @@ -18,17 +18,17 @@ from pydantic import ( from pydantic.json_schema import SkipJsonSchema from generalresearch.config import is_debug -from generalresearch.currency import USDCent from generalresearch.decorators import LOG from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.thl.definitions import SessionAdjustedStatus -from generalresearch.pg_helper import PostgresConfig payout_example = random.randint(150, 750 * 100) adjustment_example = random.randint(-1_000, 50 * 100) if TYPE_CHECKING: + from generalresearch.currency import USDCent from generalresearch.models.thl.ledger import LedgerAccount + from generalresearch.pg_helper import PostgresConfig class AdjustmentType(BaseModel): diff --git a/generalresearch/models/thl/ipinfo.py b/generalresearch/models/thl/ipinfo.py index 8fbae4c..8322c7d 100644 --- a/generalresearch/models/thl/ipinfo.py +++ b/generalresearch/models/thl/ipinfo.py @@ -15,14 +15,13 @@ from pydantic import ( field_validator, ) -from generalresearch.models.custom_types import ( - AwareDatetimeISO, - CountryISOLike, - IPvAnyAddressStr, -) - if TYPE_CHECKING: from generalresearch.managers.thl.ipinfo import IPGeonameManager + from generalresearch.models.custom_types import ( + AwareDatetimeISO, + CountryISOLike, + IPvAnyAddressStr, + ) fake = Faker() diff --git a/generalresearch/models/thl/leaderboard.py b/generalresearch/models/thl/leaderboard.py index 4d116a5..3b33fe6 100644 --- a/generalresearch/models/thl/leaderboard.py +++ b/generalresearch/models/thl/leaderboard.py @@ -4,7 +4,7 @@ import logging import math from datetime import UTC, datetime, timedelta from enum import StrEnum -from typing import Literal +from typing import TYPE_CHECKING, Literal from uuid import UUID, uuid3 from zoneinfo import ZoneInfo @@ -21,9 +21,12 @@ from pydantic import ( from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.legacy.api_status import StatusResponse -from generalresearch.models.thl.locales import CountryISO from generalresearch.utils.enum import ReprEnumMeta +if TYPE_CHECKING: + from generalresearch.models.thl.locales import CountryISO + + logger = logging.getLogger() diff --git a/generalresearch/models/thl/ledger.py b/generalresearch/models/thl/ledger.py index 518e390..c38e83b 100644 --- a/generalresearch/models/thl/ledger.py +++ b/generalresearch/models/thl/ledger.py @@ -2,7 +2,7 @@ from __future__ import annotations from datetime import UTC, datetime from enum import IntEnum, StrEnum -from typing import Annotated, Any, Literal, Self +from typing import TYPE_CHECKING, Annotated, Any, Literal, Self from uuid import uuid4 from pydantic import ( @@ -16,12 +16,7 @@ from pydantic import ( model_validator, ) -from generalresearch.models.custom_types import ( - AwareDatetimeISO, - HttpsUrlStr, - UUIDStr, - check_valid_uuid, -) +from generalresearch.models.custom_types import check_valid_uuid from generalresearch.models.thl.pagination import Page from generalresearch.models.thl.payout_format import ( PayoutFormatType, @@ -29,6 +24,16 @@ from generalresearch.models.thl.payout_format import ( ) from generalresearch.utils.enum import ReprEnumMeta +if TYPE_CHECKING: + from generalresearch.models.custom_types import ( + AwareDatetimeISO, + HttpsUrlStr, + UUIDStr, + ) + from generalresearch.models.thl.payout_format import ( + PayoutFormatType, + ) + def _example_user_tx_payout(schema: dict[str, Any]) -> None: diff --git a/generalresearch/models/thl/offerwall/base.py b/generalresearch/models/thl/offerwall/base.py index 33b9847..1d41ef2 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 Annotated, Any, Self +from typing import TYPE_CHECKING, Annotated, Any, Self from uuid import uuid4 import numpy as np @@ -20,31 +20,37 @@ from pydantic import ( ) from generalresearch.models import Source -from generalresearch.models.custom_types import HttpsUrl, UUIDStr from generalresearch.models.legacy.bucket import ( Bucket as LegacyBucket, ) from generalresearch.models.legacy.bucket import ( - CategoryAssociation, DurationSummary, - Eligibility, PayoutSummary, PayoutSummaryDecimal, - SurveyEligibilityCriterion, ) from generalresearch.models.legacy.definitions import OfferwallReason -from generalresearch.models.thl.locales import CountryISO from generalresearch.models.thl.offerwall import ( OFFERWALL_TYPE_CLASS, - OfferWallType, - OfferWallTypeClass, ) from generalresearch.models.thl.offerwall.bucket import ( generate_offerwall_entry_url, ) -from generalresearch.models.thl.profiling.upk_question import UpkQuestion from generalresearch.models.thl.soft_pair import SoftPairResultType -from generalresearch.models.thl.user import User + +if TYPE_CHECKING: + from generalresearch.models.custom_types import HttpsUrl, UUIDStr + from generalresearch.models.legacy.bucket import ( + CategoryAssociation, + Eligibility, + SurveyEligibilityCriterion, + ) + from generalresearch.models.thl.locales import CountryISO + from generalresearch.models.thl.offerwall import ( + OfferWallType, + OfferWallTypeClass, + ) + from generalresearch.models.thl.profiling.upk_question import UpkQuestion + from generalresearch.models.thl.user import User class MergeTableFeatures(BaseModel): diff --git a/generalresearch/models/thl/offerwall/cache.py b/generalresearch/models/thl/offerwall/cache.py index 97546b2..aa18014 100644 --- a/generalresearch/models/thl/offerwall/cache.py +++ b/generalresearch/models/thl/offerwall/cache.py @@ -1,18 +1,19 @@ from __future__ import annotations from datetime import UTC, datetime -from typing import Any +from typing import TYPE_CHECKING, Any from pydantic import BaseModel, Field -from generalresearch.models import Source -from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr -from generalresearch.models.thl.offerwall import OfferWallRequest -from generalresearch.models.thl.offerwall.base import ( - OfferwallBase, - ScoredTaskResult, - TaskResult, -) +if TYPE_CHECKING: + from generalresearch.models import Source + from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr + from generalresearch.models.thl.offerwall import OfferWallRequest + from generalresearch.models.thl.offerwall.base import ( + OfferwallBase, + ScoredTaskResult, + TaskResult, + ) class GetOfferWallCache(BaseModel): diff --git a/generalresearch/models/thl/payout.py b/generalresearch/models/thl/payout.py index ce8a809..9902af3 100644 --- a/generalresearch/models/thl/payout.py +++ b/generalresearch/models/thl/payout.py @@ -2,7 +2,7 @@ from __future__ import annotations import json from datetime import UTC, datetime -from typing import Self +from typing import TYPE_CHECKING, Self from uuid import uuid4 from pydantic import ( @@ -17,12 +17,18 @@ from pydantic import ( from pydantic.json_schema import SkipJsonSchema from generalresearch.currency import USDCent -from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr, UUIDStrCoerce from generalresearch.models.thl.definitions import PayoutStatus from generalresearch.models.thl.wallet import PayoutType -from generalresearch.models.thl.wallet.cashout_method import ( - CashMailOrderData, -) + +if TYPE_CHECKING: + from generalresearch.models.custom_types import ( + AwareDatetimeISO, + UUIDStr, + UUIDStrCoerce, + ) + from generalresearch.models.thl.wallet.cashout_method import ( + CashMailOrderData, + ) class PayoutEvent(BaseModel): @@ -59,9 +65,7 @@ class PayoutEvent(BaseModel): examples=["a6dc1fc1bf934557b952f253dee12813"], ) - created: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=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. @@ -233,9 +237,7 @@ class BusinessPayoutEventCreate(BaseModel): examples=[uuid4().hex], ) - created: AwareDatetimeISO = Field( - default_factory=lambda: datetime.now(tz=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. @@ -351,5 +353,6 @@ class BusinessPayoutEventCreate(BaseModel): ) return d + class BusinessPayoutEvent(BusinessPayoutEventCreate): id: SkipJsonSchema[PositiveInt] = Field(exclude=True) diff --git a/generalresearch/models/thl/product.py b/generalresearch/models/thl/product.py index 83955a6..988b72d 100644 --- a/generalresearch/models/thl/product.py +++ b/generalresearch/models/thl/product.py @@ -45,7 +45,13 @@ from generalresearch.models.custom_types import ( HttpsUrlStr, UUIDStr, ) -from generalresearch.models.thl.ledger import LedgerAccount +from generalresearch.models.thl.finance import ( + POPFinancial, + ProductBalances, +) +from generalresearch.models.thl.payout import ( + BrokerageProductPayoutEvent, +) from generalresearch.models.thl.payout_format import ( PayoutFormatType, format_payout_format, @@ -70,13 +76,7 @@ if TYPE_CHECKING: from generalresearch.managers.thl.payout import ( BrokerageProductPayoutEventManager, ) - from generalresearch.models.thl.finance import ( - POPFinancial, - ProductBalances, - ) - from generalresearch.models.thl.payout import ( - BrokerageProductPayoutEvent, - ) + from generalresearch.models.thl.ledger import LedgerAccount # fmt: off diff --git a/generalresearch/models/thl/profiling/marketplace.py b/generalresearch/models/thl/profiling/marketplace.py index 0129e38..0c1e39b 100644 --- a/generalresearch/models/thl/profiling/marketplace.py +++ b/generalresearch/models/thl/profiling/marketplace.py @@ -3,18 +3,21 @@ from __future__ import annotations from abc import ABC, abstractmethod from datetime import UTC, datetime from functools import cached_property -from typing import Any +from typing import TYPE_CHECKING, Any from pydantic import BaseModel, ConfigDict, Field, PositiveInt, computed_field -from generalresearch.models import MAX_INT32, Source -from generalresearch.models.custom_types import ( - AwareDatetimeISO, - CountryISOLike, - LanguageISOLike, - UUIDStr, -) -from generalresearch.models.thl.locales import CountryISO, LanguageISO +from generalresearch.models import MAX_INT32 + +if TYPE_CHECKING: + from generalresearch.models import Source + from generalresearch.models.custom_types import ( + AwareDatetimeISO, + CountryISOLike, + LanguageISOLike, + UUIDStr, + ) + from generalresearch.models.thl.locales import CountryISO, LanguageISO class MarketplaceQuestion(BaseModel, ABC): diff --git a/generalresearch/models/thl/profiling/question.py b/generalresearch/models/thl/profiling/question.py index 3e2984a..920dd3a 100644 --- a/generalresearch/models/thl/profiling/question.py +++ b/generalresearch/models/thl/profiling/question.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import Any +from typing import TYPE_CHECKING, Any from pydantic import ( BaseModel, @@ -9,13 +9,14 @@ from pydantic import ( computed_field, ) -from generalresearch.models.custom_types import ( - AwareDatetimeISO, - CountryISOLike, - LanguageISOLike, - UUIDStr, -) -from generalresearch.models.thl.profiling.upk_question import UpkQuestion +if TYPE_CHECKING: + from generalresearch.models.custom_types import ( + AwareDatetimeISO, + CountryISOLike, + LanguageISOLike, + UUIDStr, + ) + from generalresearch.models.thl.profiling.upk_question import UpkQuestion class Question(BaseModel): diff --git a/generalresearch/models/thl/profiling/upk_property.py b/generalresearch/models/thl/profiling/upk_property.py index e46e00a..922f5a4 100644 --- a/generalresearch/models/thl/profiling/upk_property.py +++ b/generalresearch/models/thl/profiling/upk_property.py @@ -2,14 +2,17 @@ from __future__ import annotations from enum import StrEnum from functools import cached_property +from typing import TYPE_CHECKING from uuid import uuid4 from pydantic import BaseModel, ConfigDict, Field, TypeAdapter -from generalresearch.models.custom_types import CountryISOLike, UUIDStr -from generalresearch.models.thl.category import Category from generalresearch.utils.enum import ReprEnumMeta +if TYPE_CHECKING: + from generalresearch.models.custom_types import CountryISOLike, UUIDStr + from generalresearch.models.thl.category import Category + class PropertyType(StrEnum, metaclass=ReprEnumMeta): # UserProfileKnowledge Item diff --git a/generalresearch/models/thl/profiling/upk_question.py b/generalresearch/models/thl/profiling/upk_question.py index 3bb0733..a73683c 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 Annotated, Any, Literal +from typing import TYPE_CHECKING, Annotated, Any, Literal from pydantic import ( BaseModel, @@ -18,9 +18,11 @@ from pydantic import ( ) from generalresearch.models import Source -from generalresearch.models.custom_types import UUIDStr 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 d8323ad..41895b1 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 UTC, datetime -from typing import Any, Self +from typing import TYPE_CHECKING, Any, Self from uuid import uuid4 from pydantic import ( @@ -14,16 +14,18 @@ from pydantic import ( ) from generalresearch.models import MAX_INT32 -from generalresearch.models.custom_types import ( - AwareDatetimeISO, - CountryISOLike, - UUIDStr, -) from generalresearch.models.thl.profiling.upk_property import ( Cardinality, PropertyType, ) +if TYPE_CHECKING: + from generalresearch.models.custom_types import ( + AwareDatetimeISO, + CountryISOLike, + UUIDStr, + ) + class UpkQuestionAnswer(BaseModel): diff --git a/generalresearch/models/thl/profiling/user_info.py b/generalresearch/models/thl/profiling/user_info.py index 32af704..c82e2d2 100644 --- a/generalresearch/models/thl/profiling/user_info.py +++ b/generalresearch/models/thl/profiling/user_info.py @@ -1,14 +1,17 @@ from __future__ import annotations +from typing import TYPE_CHECKING + from pydantic import BaseModel, ConfigDict, Field from pydantic.json_schema import SkipJsonSchema -from generalresearch.models import Source -from generalresearch.models.custom_types import AwareDatetimeISO -from generalresearch.models.thl.profiling.user_question_answer import ( - MarketplaceResearchProfileQuestion, -) -from generalresearch.models.thl.user import User +if TYPE_CHECKING: + from generalresearch.models import Source + from generalresearch.models.custom_types import AwareDatetimeISO + from generalresearch.models.thl.profiling.user_question_answer import ( + MarketplaceResearchProfileQuestion, + ) + from generalresearch.models.thl.user import User class UserProfileKnowledgeAnswer(BaseModel): diff --git a/generalresearch/models/thl/profiling/user_question_answer.py b/generalresearch/models/thl/profiling/user_question_answer.py index 378345e..b1868b3 100644 --- a/generalresearch/models/thl/profiling/user_question_answer.py +++ b/generalresearch/models/thl/profiling/user_question_answer.py @@ -3,7 +3,7 @@ from __future__ import annotations import json from collections.abc import Iterator from datetime import UTC, datetime, timedelta -from typing import Any, Literal +from typing import TYPE_CHECKING, Any, Literal from pydantic import ( BaseModel, @@ -14,10 +14,13 @@ from pydantic import ( model_validator, ) -from generalresearch.models import MAX_INT32, Source -from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr -from generalresearch.models.thl.locales import CountryISO, LanguageISO -from generalresearch.models.thl.profiling.upk_question import UpkQuestion +from generalresearch.models import MAX_INT32 + +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.thl.profiling.upk_question import UpkQuestion class UserQuestionAnswer(BaseModel): diff --git a/generalresearch/models/thl/report_task.py b/generalresearch/models/thl/report_task.py index 299ba90..e4a8f48 100644 --- a/generalresearch/models/thl/report_task.py +++ b/generalresearch/models/thl/report_task.py @@ -3,11 +3,14 @@ from __future__ import annotations import random from collections import defaultdict from collections.abc import Collection +from typing import TYPE_CHECKING from pydantic import BaseModel, ConfigDict, Field from generalresearch.models.thl.definitions import ReportValue -from generalresearch.models.thl.user import BPUIDStr + +if TYPE_CHECKING: + from generalresearch.models.thl.user import BPUIDStr # If a report is made with multiple values, we'll take the one with the # highest priority diff --git a/generalresearch/models/thl/session.py b/generalresearch/models/thl/session.py index 31dc668..871e5c4 100644 --- a/generalresearch/models/thl/session.py +++ b/generalresearch/models/thl/session.py @@ -18,14 +18,7 @@ from pydantic import ( model_validator, ) -from generalresearch.models import DeviceType, Source -from generalresearch.models.custom_types import ( - AwareDatetimeISO, - EnumNameSerializer, - IPvAnyAddressStr, - UUIDStr, -) -from generalresearch.models.legacy.bucket import Bucket +from generalresearch.models import Source from generalresearch.models.thl import ( decimal_to_int_cents, int_cents_to_decimal, @@ -33,9 +26,7 @@ from generalresearch.models.thl import ( from generalresearch.models.thl.definitions import ( WALL_ALLOWED_STATUS_CODE_1_2, WALL_ALLOWED_STATUS_STATUS_CODE, - ReportValue, SessionAdjustedStatus, - SessionStatusCode2, Status, StatusCode1, WallAdjustedStatus, @@ -46,6 +37,18 @@ 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.legacy.bucket import Bucket + from generalresearch.models.thl.definitions import ( + ReportValue, + SessionStatusCode2, + ) from generalresearch.models.thl.product import Product from generalresearch.models.thl.user import User diff --git a/generalresearch/models/thl/soft_pair.py b/generalresearch/models/thl/soft_pair.py index 7c2f36e..313d374 100644 --- a/generalresearch/models/thl/soft_pair.py +++ b/generalresearch/models/thl/soft_pair.py @@ -2,12 +2,14 @@ from __future__ import annotations from dataclasses import dataclass from enum import Enum +from typing import TYPE_CHECKING -from generalresearch.models import Source -from generalresearch.models.dynata.survey import DynataCondition -from generalresearch.models.thl.survey.condition import ( - MarketplaceCondition, -) +if TYPE_CHECKING: + from generalresearch.models import Source + from generalresearch.models.dynata.survey import DynataCondition + from generalresearch.models.thl.survey.condition import ( + MarketplaceCondition, + ) class SoftPairResultType(int, Enum): diff --git a/generalresearch/models/thl/survey/__init__.py b/generalresearch/models/thl/survey/__init__.py index 6e6b475..d0f2b33 100644 --- a/generalresearch/models/thl/survey/__init__.py +++ b/generalresearch/models/thl/survey/__init__.py @@ -3,27 +3,32 @@ from __future__ import annotations from abc import ABC, abstractmethod from decimal import Decimal from itertools import product +from typing import TYPE_CHECKING from more_itertools import flatten from pydantic import BaseModel, Field -from generalresearch.models import Source from generalresearch.models.thl.demographics import ( AgeGroup, DemographicTarget, Gender, ) -from generalresearch.models.thl.locales import ( - CountryISO, - CountryISOs, - LanguageISO, - LanguageISOs, -) from generalresearch.models.thl.survey.condition import ( ConditionValueType, - MarketplaceCondition, ) +if TYPE_CHECKING: + from generalresearch.models import Source + from generalresearch.models.thl.locales import ( + CountryISO, + CountryISOs, + LanguageISO, + LanguageISOs, + ) + from generalresearch.models.thl.survey.condition import ( + MarketplaceCondition, + ) + class MarketplaceTask(BaseModel, ABC): """This is called a "Task" even though generally it represents a survey diff --git a/generalresearch/models/thl/survey/buyer.py b/generalresearch/models/thl/survey/buyer.py index b888007..26846d3 100644 --- a/generalresearch/models/thl/survey/buyer.py +++ b/generalresearch/models/thl/survey/buyer.py @@ -3,7 +3,7 @@ from __future__ import annotations from datetime import UTC, datetime from decimal import Decimal from math import log -from typing import Annotated +from typing import TYPE_CHECKING, Annotated from pydantic import ( BaseModel, @@ -17,11 +17,13 @@ from pydantic import ( from scipy.stats import beta as beta_dist from generalresearch.models import Source -from generalresearch.models.custom_types import ( - AwareDatetimeISO, - CountryISOLike, - UUIDStr, -) + +if TYPE_CHECKING: + from generalresearch.models.custom_types import ( + AwareDatetimeISO, + CountryISOLike, + UUIDStr, + ) class Buyer(BaseModel): diff --git a/generalresearch/models/thl/survey/model.py b/generalresearch/models/thl/survey/model.py index 2eed8f7..9fa3d8e 100644 --- a/generalresearch/models/thl/survey/model.py +++ b/generalresearch/models/thl/survey/model.py @@ -2,7 +2,7 @@ from __future__ import annotations from datetime import UTC, datetime from decimal import Decimal -from typing import Annotated, Any +from typing import TYPE_CHECKING, Annotated, Any from pydantic import ( BaseModel, @@ -17,18 +17,21 @@ from pydantic import ( ) from generalresearch.managers.thl.buyer import Buyer -from generalresearch.models import Source -from generalresearch.models.custom_types import ( - AwareDatetimeISO, - CountryISOLike, - EnumNameSerializer, - PropertyCode, - SurveyKey, -) -from generalresearch.models.thl.category import Category -from generalresearch.models.thl.definitions import Status, StatusCode1 +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, + EnumNameSerializer, + PropertyCode, + SurveyKey, + ) + from generalresearch.models.thl.category import Category + from generalresearch.models.thl.definitions import Status + class SurveyCategoryModel(BaseModel): model_config = ConfigDict(from_attributes=True) diff --git a/generalresearch/models/thl/survey/penalty.py b/generalresearch/models/thl/survey/penalty.py index 755d25c..54edb94 100644 --- a/generalresearch/models/thl/survey/penalty.py +++ b/generalresearch/models/thl/survey/penalty.py @@ -2,15 +2,16 @@ from __future__ import annotations import abc from datetime import UTC, datetime -from typing import Annotated, Literal +from typing import TYPE_CHECKING, Annotated, Literal from pydantic import BaseModel, ConfigDict, Field, TypeAdapter -from generalresearch.models import Source -from generalresearch.models.custom_types import ( - AwareDatetimeISO, - UUIDStr, -) +if TYPE_CHECKING: + from generalresearch.models import Source + from generalresearch.models.custom_types import ( + AwareDatetimeISO, + UUIDStr, + ) class SurveyPenalty(BaseModel, abc.ABC): diff --git a/generalresearch/models/thl/survey/task_collection.py b/generalresearch/models/thl/survey/task_collection.py index d8db0d1..c9614d5 100644 --- a/generalresearch/models/thl/survey/task_collection.py +++ b/generalresearch/models/thl/survey/task_collection.py @@ -3,12 +3,14 @@ from __future__ import annotations import copy import json import logging +from typing import TYPE_CHECKING import pandas as pd import pandera.pandas as pa from pydantic import BaseModel, ConfigDict, Field, model_validator -from generalresearch.models.thl.survey import MarketplaceTask +if TYPE_CHECKING: + from generalresearch.models.thl.survey import MarketplaceTask logging.basicConfig() logger = logging.getLogger() diff --git a/generalresearch/models/thl/task_adjustment.py b/generalresearch/models/thl/task_adjustment.py index 27c47d4..fa5592e 100644 --- a/generalresearch/models/thl/task_adjustment.py +++ b/generalresearch/models/thl/task_adjustment.py @@ -2,16 +2,20 @@ from __future__ import annotations from datetime import UTC, datetime from decimal import Decimal +from typing import TYPE_CHECKING from uuid import uuid4 from pydantic import BaseModel, ConfigDict, Field, PositiveInt, model_validator -from generalresearch.models import MAX_INT32, Source -from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr +from generalresearch.models 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 + class TaskAdjustmentEvent(BaseModel): """ diff --git a/generalresearch/models/thl/task_status.py b/generalresearch/models/thl/task_status.py index 7719b18..817f4c5 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 TYPE_CHECKING, Annotated, Any, Literal from pydantic import ( BaseModel, @@ -13,11 +13,6 @@ from pydantic import ( model_validator, ) -from generalresearch.models.custom_types import ( - AwareDatetimeISO, - EnumNameSerializer, - UUIDStr, -) from generalresearch.models.thl import decimal_to_int_cents from generalresearch.models.thl.definitions import ( SessionAdjustedStatus, @@ -28,13 +23,23 @@ from generalresearch.models.thl.definitions import ( from generalresearch.models.thl.pagination import Page from generalresearch.models.thl.payout_format import ( PayoutFormatOptionalField, - PayoutFormatType, -) -from generalresearch.models.thl.product import ( - PayoutTransformation, - Product, ) -from generalresearch.models.thl.session import Session, WallOut +from generalresearch.models.thl.session import WallOut + +if TYPE_CHECKING: + from generalresearch.models.custom_types import ( + AwareDatetimeISO, + EnumNameSerializer, + UUIDStr, + ) + from generalresearch.models.thl.payout_format import ( + PayoutFormatType, + ) + from generalresearch.models.thl.product import ( + PayoutTransformation, + Product, + ) + from generalresearch.models.thl.session import Session # API uses the ints, b/c this is what the grpc returned originally ... STATUS_MAP = { @@ -171,12 +176,12 @@ class TaskStatusResponse(BaseModel): # Serialize enum → int @field_serializer("status", return_type=int) - def serialize_status(self, v: Status | None, _info): + def serialize_status(self, v: Status | None): return STATUS_MAP[v] # Accept int OR string for input, but internally store a Status enum @field_validator("status", mode="before") - def deserialize_status(cls, v): + def deserialize_status(cls, v: Status | None): # int → enum if isinstance(v, int): return REVERSE_STATUS_MAP[v] diff --git a/generalresearch/models/thl/user.py b/generalresearch/models/thl/user.py index d3ffb0d..302aa72 100644 --- a/generalresearch/models/thl/user.py +++ b/generalresearch/models/thl/user.py @@ -21,18 +21,18 @@ from pydantic import ( from sentry_sdk import set_tag, set_user from generalresearch.models import MAX_INT32 -from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr -from generalresearch.models.thl.ipinfo import GeoIPInformation -from generalresearch.models.thl.ledger import LedgerTransaction -from generalresearch.models.thl.product import Product -from generalresearch.models.thl.userhealth import AuditLog -from generalresearch.pg_helper import PostgresConfig if TYPE_CHECKING: from generalresearch.managers.thl.ledger_manager.thl_ledger import ( ThlLedgerManager, ) from generalresearch.managers.thl.userhealth import AuditLogManager + from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr + from generalresearch.models.thl.ipinfo import GeoIPInformation + from generalresearch.models.thl.ledger import LedgerTransaction + from generalresearch.models.thl.product import Product + from generalresearch.models.thl.userhealth import AuditLog + from generalresearch.pg_helper import PostgresConfig # from generalresearch.managers.thl.userhealth import UserIpHistoryManager diff --git a/generalresearch/models/thl/user_iphistory.py b/generalresearch/models/thl/user_iphistory.py index 812c18f..0d25322 100644 --- a/generalresearch/models/thl/user_iphistory.py +++ b/generalresearch/models/thl/user_iphistory.py @@ -5,7 +5,6 @@ from datetime import UTC, datetime, timedelta from typing import TYPE_CHECKING, Self from faker import Faker -from grip_client.enums import AccessType from pydantic import ( BaseModel, ConfigDict, @@ -14,18 +13,20 @@ from pydantic import ( field_validator, ) -from generalresearch.models.custom_types import ( - AwareDatetimeISO, - CountryISOLike, - IPvAnyAddressStr, -) from generalresearch.models.thl.ipinfo import normalize_ip -from generalresearch.pg_helper import PostgresConfig -from generalresearch.redis_helper import RedisConfig if TYPE_CHECKING: + from grip_client.enums import AccessType + + from generalresearch.models.custom_types import ( + AwareDatetimeISO, + CountryISOLike, + IPvAnyAddressStr, + ) from generalresearch.models.thl.ipinfo import GeoIPInformation from generalresearch.models.thl.user import User + from generalresearch.pg_helper import PostgresConfig + from generalresearch.redis_helper import RedisConfig fake = Faker() diff --git a/generalresearch/models/thl/user_profile.py b/generalresearch/models/thl/user_profile.py index 0ec605a..2dc19b7 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 Annotated, Any, Self +from typing import TYPE_CHECKING, Annotated, Any, Self from pydantic import ( BaseModel, @@ -14,9 +14,11 @@ from pydantic import ( from pydantic.json_schema import SkipJsonSchema from generalresearch.models import MAX_INT32, Source -from generalresearch.models.custom_types import UUIDStr -from generalresearch.models.thl.user import User -from generalresearch.models.thl.user_streak import UserStreak + +if TYPE_CHECKING: + from generalresearch.models.custom_types import UUIDStr + from generalresearch.models.thl.user import User + from generalresearch.models.thl.user_streak import UserStreak class UserMetadata(BaseModel): diff --git a/generalresearch/models/thl/user_quality_event.py b/generalresearch/models/thl/user_quality_event.py index cfb4ff3..5438740 100644 --- a/generalresearch/models/thl/user_quality_event.py +++ b/generalresearch/models/thl/user_quality_event.py @@ -3,16 +3,18 @@ from __future__ import annotations from datetime import UTC, datetime from decimal import Decimal from enum import StrEnum -from typing import Literal +from typing import TYPE_CHECKING, Literal from pydantic import BaseModel, Field, PositiveInt from generalresearch.models import MAX_INT32, Source -from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr -from generalresearch.models.thl.definitions import WallAdjustedStatus -from generalresearch.models.thl.user import BPUIDStr from generalresearch.utils.enum import ReprEnumMeta +if TYPE_CHECKING: + from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr + from generalresearch.models.thl.definitions import WallAdjustedStatus + from generalresearch.models.thl.user import BPUIDStr + """ Typically used internally. These affect a user's quality standing. """ diff --git a/generalresearch/models/thl/user_streak.py b/generalresearch/models/thl/user_streak.py index cc4d643..6cd853a 100644 --- a/generalresearch/models/thl/user_streak.py +++ b/generalresearch/models/thl/user_streak.py @@ -2,6 +2,7 @@ from __future__ import annotations from datetime import date, datetime, timedelta from enum import StrEnum +from typing import TYPE_CHECKING from zoneinfo import ZoneInfo import pandas as pd @@ -19,7 +20,9 @@ from pydantic.json_schema import SkipJsonSchema from generalresearch.managers.leaderboard import country_timezone from generalresearch.models import MAX_INT32 -from generalresearch.models.thl.locales import CountryISO + +if TYPE_CHECKING: + from generalresearch.models.thl.locales import CountryISO class StreakPeriod(StrEnum): diff --git a/generalresearch/models/thl/wallet/cashout_method.py b/generalresearch/models/thl/wallet/cashout_method.py index 6158afd..1db85e8 100644 --- a/generalresearch/models/thl/wallet/cashout_method.py +++ b/generalresearch/models/thl/wallet/cashout_method.py @@ -4,7 +4,7 @@ import hashlib import logging from datetime import UTC, datetime from enum import StrEnum -from typing import Any, Literal, Self +from typing import TYPE_CHECKING, Any, Literal, Self from pydantic import ( BaseModel, @@ -17,19 +17,22 @@ from pydantic import ( model_validator, ) -from generalresearch.currency import USDCent -from generalresearch.models.custom_types import ( - AwareDatetimeISO, - HttpsUrlStr, - UUIDStr, -) from generalresearch.models.legacy.api_status import StatusResponse from generalresearch.models.thl.definitions import PayoutStatus -from generalresearch.models.thl.locales import CountryISO -from generalresearch.models.thl.user import BPUIDStr, User -from generalresearch.models.thl.wallet import Currency, PayoutType +from generalresearch.models.thl.wallet import PayoutType from generalresearch.utils.enum import ReprEnumMeta +if TYPE_CHECKING: + from generalresearch.currency import USDCent + from generalresearch.models.custom_types import ( + AwareDatetimeISO, + HttpsUrlStr, + UUIDStr, + ) + from generalresearch.models.thl.locales import CountryISO + from generalresearch.models.thl.user import BPUIDStr, User + from generalresearch.models.thl.wallet import Currency + logger = logging.getLogger() example_cashout_method = { diff --git a/generalresearch/models/thl/wallet/payout.py b/generalresearch/models/thl/wallet/payout.py index d43807d..7301b31 100644 --- a/generalresearch/models/thl/wallet/payout.py +++ b/generalresearch/models/thl/wallet/payout.py @@ -3,7 +3,7 @@ from __future__ import annotations import json from collections.abc import Collection from datetime import UTC, datetime -from typing import Any +from typing import TYPE_CHECKING, Any from uuid import uuid4 from pydantic import ( @@ -15,12 +15,14 @@ from pydantic import ( ) from generalresearch.currency import USDCent -from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.thl.definitions import PayoutStatus from generalresearch.models.thl.wallet import PayoutType -from generalresearch.models.thl.wallet.cashout_method import ( - CashMailOrderData, -) + +if TYPE_CHECKING: + from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr + from generalresearch.models.thl.wallet.cashout_method import ( + CashMailOrderData, + ) class PayoutEvent(BaseModel, validate_assignment=True): diff --git a/generalresearch/models/thl/wallet/user_wallet.py b/generalresearch/models/thl/wallet/user_wallet.py index dbe66fa..625fd52 100644 --- a/generalresearch/models/thl/wallet/user_wallet.py +++ b/generalresearch/models/thl/wallet/user_wallet.py @@ -1,15 +1,20 @@ from __future__ import annotations import logging +from typing import TYPE_CHECKING from pydantic import BaseModel, ConfigDict, Field, NonNegativeInt from generalresearch.models.legacy.api_status import StatusResponse from generalresearch.models.thl.payout_format import ( PayoutFormatField, - PayoutFormatType, ) +if TYPE_CHECKING: + from generalresearch.models.thl.payout_format import ( + PayoutFormatType, + ) + logger = logging.getLogger() example_wallet_balance = { diff --git a/generalresearch/wall_status_codes/cint.py b/generalresearch/wall_status_codes/cint.py index c31c0ba..13028dc 100644 --- a/generalresearch/wall_status_codes/cint.py +++ b/generalresearch/wall_status_codes/cint.py @@ -1,8 +1,10 @@ -from typing import Any +from typing import TYPE_CHECKING, Any -from generalresearch.models.thl.definitions import Status, StatusCode1 from generalresearch.wall_status_codes import lucid +if TYPE_CHECKING: + from generalresearch.models.thl.definitions import Status, StatusCode1 + def annotate_status_code( ext_status_code_1: str, diff --git a/test_utils/conftest.py b/test_utils/conftest.py index 041caf2..f55fe11 100644 --- a/test_utils/conftest.py +++ b/test_utils/conftest.py @@ -180,7 +180,7 @@ def git_key_path( yield Path(fn) - # os.unlink(fn) + os.unlink(fn) @pytest.fixture(scope="session") diff --git a/test_utils/managers/gr/conftest.py b/test_utils/managers/gr/conftest.py index 69f3e9a..40bd7b3 100644 --- a/test_utils/managers/gr/conftest.py +++ b/test_utils/managers/gr/conftest.py @@ -1,8 +1,11 @@ from __future__ import annotations -from collections.abc import Callable +import subprocess +from collections.abc import Callable, Generator +from random import randint import pytest +import redis import redis.asyncio as redis_async from pydantic import PostgresDsn from redis import Redis @@ -19,9 +22,16 @@ from generalresearch.redis_helper import RedisConfig # === Msc === +@pytest.fixture(scope="session") +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.gr_redis) or "127.0.0.1" in str(settings.gr_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, @@ -32,10 +42,12 @@ def gr_redis(settings: GRLBaseSettings) -> Redis: @pytest.fixture def gr_redis_async(settings: GRLBaseSettings) -> redis_async.Redis: - assert "unittest" in str(settings.gr_redis) or "127.0.0.1" in str(settings.gr_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.gr_redis), + str(settings.testing_redis), decode_responses=True, socket_timeout=0.20, socket_connect_timeout=0.20, @@ -43,21 +55,39 @@ def gr_redis_async(settings: GRLBaseSettings) -> redis_async.Redis: @pytest.fixture(scope="session") -def gr_redis_config(settings: GRLBaseSettings) -> RedisConfig: - assert "unittest" in str(settings.gr_redis) or "127.0.0.1" in str(settings.gr_redis) +def gr_redis_config( + settings: GRLBaseSettings, gr_redis_config_db: str +) -> Generator[RedisConfig]: + assert "unittest" in str(settings.testing_redis) or "127.0.0.1" in str( + settings.testing_redis + ) - return RedisConfig( - dsn=settings.gr_redis, + uri = f"redis://{settings.testing_redis}/{gr_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 gr_db(django_db_factory: Callable[..., PostgresDsn]) -> PostgresConfig: _dsn = django_db_factory("gr.common") - print("DDDD:", _dsn) return PostgresConfig( dsn=_dsn, diff --git a/test_utils/models/conftest.py b/test_utils/models/conftest.py index 5570b40..3a10ea3 100644 --- a/test_utils/models/conftest.py +++ b/test_utils/models/conftest.py @@ -22,12 +22,6 @@ from generalresearch.pg_helper import PostgresConfig if TYPE_CHECKING: from generalresearch.currency import USDCent - from generalresearch.managers.gr.business import ( - BusinessAddressManager, - BusinessBankAccountManager, - BusinessManager, - ) - from generalresearch.managers.gr.team import TeamManager from generalresearch.managers.thl.buyer import BuyerManager from generalresearch.managers.thl.ipinfo import ( IPGeonameManager, @@ -45,8 +39,6 @@ if TYPE_CHECKING: from generalresearch.managers.thl.wall import WallManager from generalresearch.models.gr.business import ( Business, - BusinessAddress, - BusinessBankAccount, ) from generalresearch.models.gr.team import Team from generalresearch.models.thl.ipinfo import IPGeoname, IPInformation @@ -407,35 +399,6 @@ def bp_payout_factory( return _inner -# === GR === - - -@pytest.fixture -def business(request, business_manager: BusinessManager) -> Business: - return business_manager.create_dummy() - - -@pytest.fixture -def business_address( - request, business: Business, business_address_manager: BusinessAddressManager -) -> BusinessAddress: - return business_address_manager.create_dummy(business_id=business.id) - - -@pytest.fixture -def business_bank_account( - request, - business: Business, - business_bank_account_manager: BusinessBankAccountManager, -) -> BusinessBankAccount: - return business_bank_account_manager.create_dummy(business_id=business.id) - - -@pytest.fixture -def team(request, team_manager: TeamManager) -> Team: - return team_manager.create_dummy() - - @pytest.fixture def audit_log(audit_log_manager: AuditLogManager, user: User) -> AuditLog: diff --git a/test_utils/models/gr/conftest.py b/test_utils/models/gr/conftest.py index 90b86aa..b623255 100644 --- a/test_utils/models/gr/conftest.py +++ b/test_utils/models/gr/conftest.py @@ -1,6 +1,7 @@ from __future__ import annotations from collections.abc import Callable +from random import randint from uuid import uuid4 import pytest @@ -20,7 +21,6 @@ from generalresearch.models.gr.business import ( Business, BusinessAddress, BusinessBankAccount, - BusinessType, TransferMethod, ) from generalresearch.models.gr.team import Membership, Team @@ -134,27 +134,38 @@ def gr_business_address_factory( @pytest.fixture def gr_business_factory( - gr_bm: BusinessManager, + gr_business_manager: BusinessManager, ) -> Callable[..., Business]: def _inner( - uuid: UUIDStr | None = None, - name: str | None = None, - team: Team | None = None, - kind: BusinessType | None = None, - tax_number: str | None = None, + save: bool = True, name: str | None = None, team: Team | None = None, **kwargs ) -> Business: - from random import randint + name = name or f"" + tax_number = str(randint(1, 999_999_999)) + + if save: + return gr_business_manager.create( + name=name, + kind="c", + uuid=uuid4().hex, + team=team, + tax_number=tax_number, + **kwargs, + ) + else: + raise ValueError("Unsaved Business not supported yet") - uuid = uuid or uuid4().hex - name = name or "< Unknown >" - tax_number = tax_number or str(randint(1, 999_999_999)) + return _inner - return gr_bm.create( - uuid=uuid, name=name, team=team, kind=kind, tax_number=tax_number - ) - return _inner +@pytest.fixture +def gr_business(gr_business_factory: Callable[..., Business]) -> Business: + return gr_business_factory(save=True) + + +@pytest.fixture +def unsaved_gr_business(gr_business_factory: Callable[..., Business]) -> Business: + return gr_business_factory(save=False) @pytest.fixture @@ -183,6 +194,21 @@ def gr_user_token( return res +@pytest.fixture +def business_address( + gr_business: Business, business_address_manager: BusinessAddressManager +) -> BusinessAddress: + return business_address_manager.create_dummy(business_id=gr_business.id) + + +@pytest.fixture +def business_bank_account( + gr_business: Business, + business_bank_account_manager: BusinessBankAccountManager, +) -> BusinessBankAccount: + return business_bank_account_manager.create_dummy(business_id=gr_business.id) + + @pytest.fixture() def gr_user_token_header(gr_user_token: GRToken) -> dict[str, str]: return gr_user_token.auth_header diff --git a/tests/models/gr/test_business.py b/tests/models/gr/test_business.py index f0107de..9310d2c 100644 --- a/tests/models/gr/test_business.py +++ b/tests/models/gr/test_business.py @@ -5,6 +5,7 @@ from collections.abc import Callable from datetime import UTC, datetime, timedelta from decimal import Decimal from pathlib import Path +from typing import TYPE_CHECKING from uuid import uuid4 import pandas as pd @@ -43,18 +44,22 @@ from generalresearch.models.thl.finance import ( BusinessBalances, ProductBalances, ) -from generalresearch.models.thl.product import BrokerageProductPayoutEvent, Product -from generalresearch.models.thl.session import Session +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.managers.thl.product import ProductManager + from generalresearch.models.thl.product import BrokerageProductPayoutEvent + from generalresearch.models.thl.session import Session + class TestBusinessBankAccount: def test_init( self, - business: Business, + gr_business: Business, business_bank_account_manager: BusinessBankAccountManager, ): from generalresearch.models.gr.business import ( @@ -63,7 +68,7 @@ class TestBusinessBankAccount: ) instance = business_bank_account_manager.create( - business_id=business.id, + business_id=gr_business.id, uuid=uuid4().hex, transfer_method=TransferMethod.ACH, ) @@ -72,7 +77,7 @@ class TestBusinessBankAccount: def test_business( self, business_bank_account: BusinessBankAccount, - business: Business, + gr_business: Business, gr_db: PostgresConfig, gr_redis_config: RedisConfig, ): @@ -84,7 +89,7 @@ class TestBusinessBankAccount: pg_config=gr_db, redis_config=gr_redis_config ) assert isinstance(business_bank_account.business, Business) - assert business_bank_account.business.uuid == business.uuid + assert business_bank_account.business.uuid == gr_business.uuid class TestBusinessAddress: @@ -122,11 +127,12 @@ class TestBusiness: def test_str_and_repr( self, - business: Business, + gr_business: Business, product_factory: Callable[..., Product], thl_web_rr: PostgresConfig, ledger_manager: LedgerManager, thl_ledger_manager: ThlLedgerManager, + product_manager: ProductManager, business_payout_event_manager: BusinessPayoutEventManager, bp_payout_factory: Callable[..., BusinessPayoutEventManager], start: datetime, @@ -139,28 +145,28 @@ class TestBusiness: create_main_accounts: Callable[..., None], ): create_main_accounts() - p1 = product_factory(business=business) + p1 = product_factory(business=gr_business) u1 = user_factory(product=p1) - p2 = product_factory(business=business) + p2 = product_factory(business=gr_business) thl_ledger_manager.get_account_or_create_bp_wallet(product=p1) thl_ledger_manager.get_account_or_create_bp_wallet(product=p2) - res1 = repr(business) + res1 = repr(gr_business) - assert business.uuid in res1 + assert gr_business.uuid in res1 assert " 0 ] ) - assert business.balance.retainer == approx(predicted_retainer, rel=0.01) + assert gr_business.balance.retainer == approx(predicted_retainer, rel=0.01) def test_neg_balance_cache( self, @@ -837,7 +845,7 @@ class TestBusinessBalance: create_main_accounts: Callable[..., None], delete_df_collection: Callable[..., None], ledger_collection, - business: Business, + gr_business: Business, user_factory: Callable[..., User], product_factory: Callable[..., Product], session_with_tx_factory: Callable[..., Session], @@ -858,8 +866,8 @@ class TestBusinessBalance: create_main_accounts() delete_df_collection(coll=ledger_collection) - p1: Product = product_factory(business=business) - p2: Product = product_factory(business=business) + p1: Product = product_factory(business=gr_business) + p2: Product = product_factory(business=gr_business) u1: User = user_factory(product=p1) u2: User = user_factory(product=p2) thl_ledger_manager.get_account_or_create_bp_wallet(product=p1) @@ -901,7 +909,7 @@ class TestBusinessBalance: 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, @@ -910,8 +918,8 @@ class TestBusinessBalance: ) # Check Product 1 - assert isinstance(business.balance, BusinessBalances) - pb1 = business.balance.product_balances[0] + assert isinstance(gr_business.balance, BusinessBalances) + pb1 = gr_business.balance.product_balances[0] assert pb1.product_id == p1.uuid assert pb1.payout == 71 assert pb1.adjustment == -71 @@ -921,7 +929,7 @@ class TestBusinessBalance: assert pb1.available_balance == 0 # Check Product 2 - pb2 = business.balance.product_balances[1] + pb2 = gr_business.balance.product_balances[1] assert pb2.product_id == p2.uuid assert pb2.payout == 71 * 2 assert pb2.adjustment == 0 @@ -931,7 +939,7 @@ class TestBusinessBalance: assert pb2.available_balance == 107 # Check Business - bb1 = business.balance + bb1 = gr_business.balance assert isinstance(bb1, BusinessBalances) assert bb1.payout == (71 * 3) # Raw total of completes assert bb1.adjustment == -71 # 1 Complete >> Failure @@ -950,7 +958,7 @@ class TestBusinessBalance: def test_multi_product_multi_payout_adjustment_at_timestamp( self, - business: Business, + gr_business: Business, product_factory: Callable[..., Product], user_factory: Callable[..., User], mnt_filepath: GRLDatasets, @@ -1008,9 +1016,9 @@ class TestBusinessBalance: delete_df_collection(coll=ledger_collection) delete_df_collection(coll=task_adj_collection) - u1: User = user_factory(product=product_factory(business=business)) - u2: User = user_factory(product=product_factory(business=business)) - u3: User = user_factory(product=product_factory(business=business)) + u1: User = user_factory(product=product_factory(business=gr_business)) + u2: User = user_factory(product=product_factory(business=gr_business)) + u3: User = user_factory(product=product_factory(business=gr_business)) s1 = session_with_tx_factory( user=u1, @@ -1063,7 +1071,7 @@ class TestBusinessBalance: df = client_no_amm.compute(pop_ledger_merge.ddf(), sync=True) assert df.shape == (20, 28) - business.prebuild_balance( + gr_business.prebuild_balance( thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, @@ -1071,7 +1079,7 @@ class TestBusinessBalance: pop_ledger=pop_ledger_merge, ) - business.prebuild_balance( + gr_business.prebuild_balance( thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, @@ -1079,9 +1087,9 @@ class TestBusinessBalance: pop_ledger=pop_ledger_merge, at_timestamp=start + timedelta(days=1, hours=1), ) - day1_bal = business.balance + day1_bal = gr_business.balance - business.prebuild_balance( + gr_business.prebuild_balance( thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, @@ -1089,9 +1097,9 @@ class TestBusinessBalance: pop_ledger=pop_ledger_merge, at_timestamp=start + timedelta(days=2, hours=1), ) - day2_bal = business.balance + day2_bal = gr_business.balance - business.prebuild_balance( + gr_business.prebuild_balance( thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, @@ -1099,9 +1107,9 @@ class TestBusinessBalance: pop_ledger=pop_ledger_merge, at_timestamp=start + timedelta(days=3, hours=1), ) - day3_bal = business.balance + day3_bal = gr_business.balance - business.prebuild_balance( + gr_business.prebuild_balance( thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, @@ -1109,9 +1117,9 @@ class TestBusinessBalance: pop_ledger=pop_ledger_merge, at_timestamp=start + timedelta(days=4, hours=1), ) - day4_bal = business.balance + day4_bal = gr_business.balance - business.prebuild_balance( + gr_business.prebuild_balance( thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, @@ -1119,9 +1127,9 @@ class TestBusinessBalance: pop_ledger=pop_ledger_merge, at_timestamp=start + timedelta(days=5, hours=1), ) - day5_bal = business.balance + day5_bal = gr_business.balance - business.prebuild_balance( + gr_business.prebuild_balance( thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, @@ -1129,7 +1137,7 @@ class TestBusinessBalance: pop_ledger=pop_ledger_merge, at_timestamp=start + timedelta(days=6, hours=1), ) - day6_bal = business.balance + day6_bal = gr_business.balance assert isinstance(day1_bal, BusinessBalances) assert isinstance(day2_bal, BusinessBalances) @@ -1187,7 +1195,7 @@ class TestBusinessMethods: def test_set_cache( self, - business: Business, + gr_business: Business, gr_redis: RedisConfig, gr_db: PostgresConfig, thl_web_rr: PostgresConfig, @@ -1208,9 +1216,9 @@ class TestBusinessMethods: gr_redis_config: RedisConfig, mnt_gr_api_dir: Path, ): - assert gr_redis.get(name=business.cache_key) is None + assert gr_redis.get(name=gr_business.cache_key) is None - p1 = product_factory(team=team, business=business) + p1 = product_factory(team=team, business=gr_business) u1 = user_factory(product=p1) # Business needs tx & incite to build balance @@ -1221,7 +1229,7 @@ class TestBusinessMethods: ledger_collection.initial_load(client=None, sync=True) pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) - business.set_cache( + gr_business.set_cache( pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config, @@ -1234,14 +1242,14 @@ class TestBusinessMethods: mnt_gr_api=mnt_gr_api_dir, ) - assert gr_redis.hgetall(name=business.cache_key) is not None + assert gr_redis.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 # it's being assigned to the field regardless year = datetime.now(tz=UTC).year res = Business.from_redis( - uuid=business.uuid, + uuid=gr_business.uuid, fields=[f"pop_financial:{year}"], gr_redis_config=gr_redis_config, ) @@ -1249,7 +1257,7 @@ class TestBusinessMethods: def test_set_cache_business( self, - business: Business, + gr_business: Business, gr_db: PostgresConfig, thl_web_rr: PostgresConfig, product_factory: Callable[..., Product], @@ -1272,9 +1280,9 @@ class TestBusinessMethods: ): from generalresearch.models.gr.business import Business - p1 = product_factory(team=team, business=business) + p1 = product_factory(team=team, business=gr_business) u1 = user_factory(product=p1) - team_manager.add_business(team=team, business=business) + team_manager.add_business(team=team, business=gr_business) # Business needs tx & incite to build balance delete_ledger_db() @@ -1284,7 +1292,7 @@ class TestBusinessMethods: ledger_collection.initial_load(client=None, sync=True) pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) - business.set_cache( + gr_business.set_cache( pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config, @@ -1299,7 +1307,7 @@ class TestBusinessMethods: # keys: List = Business.required_fields() + ["products", "bp_accounts"] business2 = Business.from_redis( - uuid=business.uuid, + uuid=gr_business.uuid, fields=[ "id", "tax_number", @@ -1319,7 +1327,7 @@ class TestBusinessMethods: ) assert isinstance(business2, Business) - assert business.model_dump_json() == business2.model_dump_json() + assert gr_business.model_dump_json() == business2.model_dump_json() # assert isinstance(business2.balance, BusinessBalances) assert isinstance(business2.products, list) assert isinstance(business2.teams, list) @@ -1413,7 +1421,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, ): @@ -1421,8 +1429,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) @@ -1443,7 +1451,7 @@ class TestBusinessMethods: pg_config=thl_web_rr, ) - business.prebuild_enriched_wall_parquet( + gr_business.prebuild_enriched_wall_parquet( thl_pg_config=thl_web_rr, ds=mnt_filepath, client=client_no_amm, @@ -1453,6 +1461,6 @@ class TestBusinessMethods: # Now try to read from path df = pd.read_parquet( - os.path.join(mnt_gr_api_dir, "pop_event", f"{business.file_key}.parquet") + os.path.join(mnt_gr_api_dir, "pop_event", f"{gr_business.file_key}.parquet") ) assert isinstance(df, pd.DataFrame) -- 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') 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') 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 186f5e47536673fde05aa6c3e0025f915e77e6a0 Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Tue, 1 Sep 2026 15:36:12 -0700 Subject: Circular resolved. Working on GR Model tests --- generalresearch/grliq/managers/forensic_data.py | 3 +- generalresearch/grliq/managers/forensic_events.py | 3 +- generalresearch/grliq/utils.py | 1 - generalresearch/managers/events.py | 2 +- generalresearch/managers/gr/authentication.py | 3 +- generalresearch/managers/gr/business.py | 3 +- generalresearch/managers/gr/team.py | 3 +- generalresearch/managers/network/label.py | 5 ++- generalresearch/managers/thl/category.py | 2 +- generalresearch/managers/thl/contest_manager.py | 2 +- .../managers/thl/ledger_manager/conditions.py | 2 +- .../managers/thl/ledger_manager/ledger.py | 4 +-- .../managers/thl/ledger_manager/thl_ledger.py | 3 +- generalresearch/managers/thl/payout.py | 2 +- generalresearch/managers/thl/product.py | 3 +- generalresearch/managers/thl/session.py | 3 +- generalresearch/managers/thl/survey_penalty.py | 5 ++- generalresearch/managers/thl/task_adjustment.py | 3 +- generalresearch/managers/thl/user_compensate.py | 3 +- .../thl/user_manager/mysql_user_manager.py | 3 +- .../managers/thl/user_manager/user_manager.py | 2 +- generalresearch/managers/thl/wall.py | 2 +- generalresearch/models/admin/request.py | 5 ++- generalresearch/models/cint/question.py | 2 +- generalresearch/models/cint/survey.py | 11 +++--- generalresearch/models/dynata/question.py | 6 ++-- generalresearch/models/dynata/survey.py | 15 ++++---- generalresearch/models/events.py | 12 ++++--- generalresearch/models/gr/__init__.py | 22 ++++++------ generalresearch/models/gr/authentication.py | 6 ++-- generalresearch/models/gr/business.py | 5 ++- generalresearch/models/gr/team.py | 10 +++--- generalresearch/models/innovate/survey.py | 13 +++---- generalresearch/models/legacy/bucket.py | 14 ++++---- generalresearch/models/legacy/offerwall.py | 3 +- generalresearch/models/legacy/questions.py | 2 +- generalresearch/models/lucid/survey.py | 12 +++---- generalresearch/models/morning/survey.py | 9 ++--- generalresearch/models/network/mtr/execute.py | 5 +-- generalresearch/models/network/nmap/execute.py | 5 +-- generalresearch/models/network/rdns/execute.py | 5 +-- generalresearch/models/network/tool_run.py | 12 ++++--- generalresearch/models/precision/survey.py | 14 ++++---- generalresearch/models/prodege/survey.py | 15 ++++---- generalresearch/models/repdata/question.py | 2 +- generalresearch/models/repdata/survey.py | 14 ++++---- generalresearch/models/spectrum/survey.py | 18 +++++----- generalresearch/models/thl/__init__.py | 40 ++++++++++++---------- generalresearch/models/thl/contest/__init__.py | 2 +- generalresearch/models/thl/contest/contest.py | 2 +- .../models/thl/contest/contest_entry.py | 2 +- generalresearch/models/thl/contest/milestone.py | 6 ++-- generalresearch/models/thl/finance.py | 23 ++----------- generalresearch/models/thl/ledger.py | 12 +++---- generalresearch/models/thl/offerwall/base.py | 2 +- generalresearch/models/thl/offerwall/cache.py | 3 +- generalresearch/models/thl/payout.py | 18 +++++----- .../models/thl/profiling/marketplace.py | 13 +++---- generalresearch/models/thl/profiling/question.py | 13 +++---- .../models/thl/profiling/upk_property.py | 2 +- .../models/thl/profiling/upk_question_answer.py | 14 ++++---- generalresearch/models/thl/profiling/user_info.py | 3 +- generalresearch/models/thl/session.py | 12 +++---- generalresearch/models/thl/survey/buyer.py | 14 ++++---- generalresearch/models/thl/survey/model.py | 15 ++++---- generalresearch/models/thl/survey/penalty.py | 10 +++--- generalresearch/models/thl/task_adjustment.py | 2 +- generalresearch/models/thl/task_status.py | 11 +++--- generalresearch/models/thl/user.py | 2 +- generalresearch/models/thl/user_profile.py | 3 +- generalresearch/models/thl/user_quality_event.py | 2 +- .../models/thl/wallet/cashout_method.py | 13 +++---- generalresearch/models/thl/wallet/payout.py | 12 +++---- test_utils/models/gr/conftest.py | 3 +- test_utils/models/thl/conftest.py | 10 +++--- tests/managers/thl/test_ledger/test_lm_accounts.py | 4 +-- tests/models/custom_types/test_aware_datetime.py | 4 +-- tests/models/custom_types/test_uuid_str.py | 4 +-- tests/models/thl/test_product.py | 4 +-- 79 files changed, 278 insertions(+), 301 deletions(-) (limited to 'test_utils/models') diff --git a/generalresearch/grliq/managers/forensic_data.py b/generalresearch/grliq/managers/forensic_data.py index 0f8534c..f523b6e 100644 --- a/generalresearch/grliq/managers/forensic_data.py +++ b/generalresearch/grliq/managers/forensic_data.py @@ -14,9 +14,10 @@ from generalresearch.grliq.models.forensic_result import ( GrlIqForensicCategoryResult, Phase, ) +from generalresearch.models.custom_types import UUIDStr if TYPE_CHECKING: - from generalresearch.models.custom_types import UUIDStr + from generalresearch.models.thl.user import User from generalresearch.pg_helper import PostgresConfig diff --git a/generalresearch/grliq/managers/forensic_events.py b/generalresearch/grliq/managers/forensic_events.py index a97a9c2..fc1ae3e 100644 --- a/generalresearch/grliq/managers/forensic_events.py +++ b/generalresearch/grliq/managers/forensic_events.py @@ -14,9 +14,10 @@ from generalresearch.grliq.models.events import ( PointerMove, TimingData, ) +from generalresearch.models.custom_types import UUIDStr if TYPE_CHECKING: - from generalresearch.models.custom_types import UUIDStr + from generalresearch.pg_helper import PostgresConfig diff --git a/generalresearch/grliq/utils.py b/generalresearch/grliq/utils.py index bceaa30..711e562 100644 --- a/generalresearch/grliq/utils.py +++ b/generalresearch/grliq/utils.py @@ -5,7 +5,6 @@ from datetime import UTC, datetime from pathlib import Path from uuid import UUID -# from generalresearch.config import from generalresearch.models.custom_types import UUIDStr diff --git a/generalresearch/managers/events.py b/generalresearch/managers/events.py index 30cec0c..4a2afb2 100644 --- a/generalresearch/managers/events.py +++ b/generalresearch/managers/events.py @@ -12,6 +12,7 @@ from redis.client import PubSub, Redis from generalresearch.incite.base import LOG from generalresearch.managers.base import RedisManager +from generalresearch.models.custom_types import UUIDStr from generalresearch.models.definitions import Source from generalresearch.models.events import ( AggregateBySource, @@ -32,7 +33,6 @@ from generalresearch.models.thl.definitions import Status if TYPE_CHECKING: from influxdb import InfluxDBClient - from generalresearch.models.custom_types import UUIDStr from generalresearch.models.events import ServerToClientMessage from generalresearch.models.thl.session import Session, Wall from generalresearch.models.thl.user import User diff --git a/generalresearch/managers/gr/authentication.py b/generalresearch/managers/gr/authentication.py index ca56467..f1ac2de 100644 --- a/generalresearch/managers/gr/authentication.py +++ b/generalresearch/managers/gr/authentication.py @@ -10,9 +10,10 @@ from psycopg import sql from pydantic import AnyHttpUrl, PositiveInt from generalresearch.managers.base import PostgresManager, PostgresManagerWithRedis +from generalresearch.models.custom_types import UUIDStr if TYPE_CHECKING: - from generalresearch.models.custom_types import UUIDStr + from generalresearch.models.gr.authentication import GRToken, GRUser from generalresearch.pg_helper import PostgresConfig from generalresearch.redis_helper import RedisConfig diff --git a/generalresearch/managers/gr/business.py b/generalresearch/managers/gr/business.py index 9bf6ef2..9338f03 100644 --- a/generalresearch/managers/gr/business.py +++ b/generalresearch/managers/gr/business.py @@ -11,6 +11,7 @@ from generalresearch.managers.base import ( PostgresManager, PostgresManagerWithRedis, ) +from generalresearch.models.custom_types import UUIDStr from generalresearch.models.gr.business import ( Business, BusinessBankAccount, @@ -18,7 +19,7 @@ from generalresearch.models.gr.business import ( 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, ) diff --git a/generalresearch/managers/gr/team.py b/generalresearch/managers/gr/team.py index 3283467..e551f85 100644 --- a/generalresearch/managers/gr/team.py +++ b/generalresearch/managers/gr/team.py @@ -11,13 +11,14 @@ from generalresearch.managers.base import ( PostgresManager, PostgresManagerWithRedis, ) +from generalresearch.models.custom_types import UUIDStr from generalresearch.models.gr.team import ( Membership, MembershipPrivilege, ) if TYPE_CHECKING: - from generalresearch.models.custom_types import UUIDStr + from generalresearch.models.gr.authentication import GRUser from generalresearch.models.gr.business import Business from generalresearch.models.gr.team import Team diff --git a/generalresearch/managers/network/label.py b/generalresearch/managers/network/label.py index aed5ff6..cdea016 100644 --- a/generalresearch/managers/network/label.py +++ b/generalresearch/managers/network/label.py @@ -9,6 +9,7 @@ from pydantic import TypeAdapter from generalresearch.managers.base import PostgresManager from generalresearch.models.custom_types import ( + AwareDatetimeISO, IPvAnyAddressStr, IPvAnyNetwork, IPvAnyNetworkStr, @@ -16,9 +17,7 @@ from generalresearch.models.custom_types import ( from generalresearch.models.network.label import IPLabel if TYPE_CHECKING: - from generalresearch.models.custom_types import ( - AwareDatetimeISO, - ) + from generalresearch.models.network.label import IPLabelKind, IPLabelSource diff --git a/generalresearch/managers/thl/category.py b/generalresearch/managers/thl/category.py index e6a091b..67c812a 100644 --- a/generalresearch/managers/thl/category.py +++ b/generalresearch/managers/thl/category.py @@ -4,11 +4,11 @@ from collections.abc import Collection from typing import TYPE_CHECKING from generalresearch.managers.base import PostgresManager +from generalresearch.models.custom_types import UUIDStr from generalresearch.models.thl.category import Category if TYPE_CHECKING: from generalresearch.managers.base import Permission - from generalresearch.models.custom_types import UUIDStr from generalresearch.pg_helper import PostgresConfig diff --git a/generalresearch/managers/thl/contest_manager.py b/generalresearch/managers/thl/contest_manager.py index 68b2cf0..b0aa505 100644 --- a/generalresearch/managers/thl/contest_manager.py +++ b/generalresearch/managers/thl/contest_manager.py @@ -10,6 +10,7 @@ from pydantic import NonNegativeInt, PositiveInt from redis import Redis from generalresearch.managers.base import PostgresManager +from generalresearch.models.custom_types import UUIDStr from generalresearch.models.thl.contest import ( ContestPrize, ContestWinner, @@ -49,7 +50,6 @@ if TYPE_CHECKING: from generalresearch.managers.thl.user_manager.user_manager import ( UserManager, ) - from generalresearch.models.custom_types import UUIDStr from generalresearch.models.thl.contest.contest import ( Contest, ContestUserView, diff --git a/generalresearch/managers/thl/ledger_manager/conditions.py b/generalresearch/managers/thl/ledger_manager/conditions.py index 7dd3021..38398b1 100644 --- a/generalresearch/managers/thl/ledger_manager/conditions.py +++ b/generalresearch/managers/thl/ledger_manager/conditions.py @@ -7,6 +7,7 @@ from typing import TYPE_CHECKING from generalresearch.config import JAMES_BILLINGS_BPID, JAMES_BILLINGS_TX_CUTOFF from generalresearch.currency import USDCent +from generalresearch.models.custom_types import UUIDStr if TYPE_CHECKING: @@ -16,7 +17,6 @@ if TYPE_CHECKING: from generalresearch.managers.thl.ledger_manager.thl_ledger import ( ThlLedgerManager, ) - from generalresearch.models.custom_types import UUIDStr from generalresearch.models.thl.product import Product from generalresearch.models.thl.session import Session, Wall from generalresearch.models.thl.user import User diff --git a/generalresearch/managers/thl/ledger_manager/ledger.py b/generalresearch/managers/thl/ledger_manager/ledger.py index f2455d4..bbf7bfa 100644 --- a/generalresearch/managers/thl/ledger_manager/ledger.py +++ b/generalresearch/managers/thl/ledger_manager/ledger.py @@ -28,7 +28,7 @@ from generalresearch.managers.thl.ledger_manager.exceptions import ( LedgerTransactionReleaseLockError, ) from generalresearch.managers.utils import parse_order_by -from generalresearch.models.custom_types import check_valid_uuid +from generalresearch.models.custom_types import UUIDStr, check_valid_uuid from generalresearch.models.thl.ledger import ( LedgerAccount, LedgerEntry, @@ -37,7 +37,7 @@ from generalresearch.models.thl.ledger import ( ) if TYPE_CHECKING: - from generalresearch.models.custom_types import UUIDStr + from generalresearch.models.thl.ledger import UserLedgerTransactionType from generalresearch.pg_helper import PostgresConfig from generalresearch.redis_helper import RedisConfig diff --git a/generalresearch/managers/thl/ledger_manager/thl_ledger.py b/generalresearch/managers/thl/ledger_manager/thl_ledger.py index bd27acf..0f8c6d3 100644 --- a/generalresearch/managers/thl/ledger_manager/thl_ledger.py +++ b/generalresearch/managers/thl/ledger_manager/thl_ledger.py @@ -29,6 +29,7 @@ from generalresearch.managers.thl.ledger_manager.conditions import ( from generalresearch.managers.thl.ledger_manager.ledger import ( LedgerManager, ) +from generalresearch.models.custom_types import UUIDStr from generalresearch.models.thl.contest.definitions import ( ContestPrizeKind, ContestType, @@ -55,7 +56,7 @@ from generalresearch.models.thl.session import Status from generalresearch.models.thl.wallet.definitions import PayoutType if TYPE_CHECKING: - from generalresearch.models.custom_types import UUIDStr + from generalresearch.models.thl.contest.contest import Contest, ContestWinner from generalresearch.models.thl.contest.raffle import ( ContestEntry, diff --git a/generalresearch/managers/thl/payout.py b/generalresearch/managers/thl/payout.py index 1749783..5397ff4 100644 --- a/generalresearch/managers/thl/payout.py +++ b/generalresearch/managers/thl/payout.py @@ -19,6 +19,7 @@ from generalresearch.managers.base import ( from generalresearch.managers.thl.ledger_manager.exceptions import ( LedgerTransactionConditionFailedError, ) +from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.thl.definitions import PayoutStatus from generalresearch.models.thl.ledger import ( Direction, @@ -42,7 +43,6 @@ if TYPE_CHECKING: ThlLedgerManager, ) from generalresearch.managers.thl.product import ProductManager - from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.gr.business import Business from generalresearch.models.thl.ledger import LedgerAccount from generalresearch.models.thl.product import Product diff --git a/generalresearch/managers/thl/product.py b/generalresearch/managers/thl/product.py index 535e566..66dd131 100644 --- a/generalresearch/managers/thl/product.py +++ b/generalresearch/managers/thl/product.py @@ -20,13 +20,12 @@ from generalresearch.decorators import LOG from generalresearch.managers.base import ( PostgresManager, ) -from generalresearch.models.custom_types import is_valid_uuid +from generalresearch.models.custom_types import UUIDStr, is_valid_uuid if TYPE_CHECKING: from generalresearch.managers.base import ( Permission, ) - from generalresearch.models.custom_types import UUIDStr from generalresearch.models.thl.product import ( PayoutConfig, Product, diff --git a/generalresearch/managers/thl/session.py b/generalresearch/managers/thl/session.py index 41d3893..5007f43 100644 --- a/generalresearch/managers/thl/session.py +++ b/generalresearch/managers/thl/session.py @@ -16,6 +16,7 @@ from generalresearch.managers.base import ( ) from generalresearch.managers.thl.product import ProductManager from generalresearch.managers.utils import parse_order_by +from generalresearch.models.custom_types import UUIDStr from generalresearch.models.legacy.bucket import Bucket from generalresearch.models.thl.session import ( Session, @@ -28,7 +29,7 @@ from generalresearch.models.thl.task_status import ( from generalresearch.models.thl.user import User if TYPE_CHECKING: - from generalresearch.models.custom_types import UUIDStr + from generalresearch.models.definitions import DeviceType from generalresearch.models.thl.definitions import ( SessionStatusCode2, diff --git a/generalresearch/managers/thl/survey_penalty.py b/generalresearch/managers/thl/survey_penalty.py index 08f8649..f176e8d 100644 --- a/generalresearch/managers/thl/survey_penalty.py +++ b/generalresearch/managers/thl/survey_penalty.py @@ -10,12 +10,11 @@ from cachetools import TTLCache, cachedmethod from generalresearch.decorators import LOG from generalresearch.managers.base import RedisManager +from generalresearch.models.custom_types import UUIDStr from generalresearch.models.thl.survey.penalty import PenaltyListAdapter if TYPE_CHECKING: - from generalresearch.models.custom_types import ( - UUIDStr, - ) + from generalresearch.models.thl.survey.penalty import ( BPSurveyPenalty, Penalty, diff --git a/generalresearch/managers/thl/task_adjustment.py b/generalresearch/managers/thl/task_adjustment.py index d0d83cb..91d914f 100644 --- a/generalresearch/managers/thl/task_adjustment.py +++ b/generalresearch/managers/thl/task_adjustment.py @@ -12,6 +12,7 @@ from generalresearch.managers.base import ( 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.custom_types import UUIDStr from generalresearch.models.thl.definitions import ( Status, WallAdjustedStatus, @@ -25,7 +26,7 @@ if TYPE_CHECKING: from generalresearch.managers.thl.ledger_manager.thl_ledger import ( ThlLedgerManager, ) - from generalresearch.models.custom_types import UUIDStr + logging.basicConfig() logger = logging.getLogger(__name__) diff --git a/generalresearch/managers/thl/user_compensate.py b/generalresearch/managers/thl/user_compensate.py index 3f0f3ec..a3cad90 100644 --- a/generalresearch/managers/thl/user_compensate.py +++ b/generalresearch/managers/thl/user_compensate.py @@ -7,11 +7,12 @@ from uuid import uuid4 from pydantic import NonNegativeInt +from generalresearch.models.custom_types import UUIDStr + if TYPE_CHECKING: from generalresearch.managers.thl.ledger_manager.thl_ledger import ( ThlLedgerManager, ) - from generalresearch.models.custom_types import UUIDStr from generalresearch.models.thl.user import User diff --git a/generalresearch/managers/thl/user_manager/mysql_user_manager.py b/generalresearch/managers/thl/user_manager/mysql_user_manager.py index 2af23dc..0b4b8a4 100644 --- a/generalresearch/managers/thl/user_manager/mysql_user_manager.py +++ b/generalresearch/managers/thl/user_manager/mysql_user_manager.py @@ -10,10 +10,11 @@ from uuid import uuid4 import psycopg from psycopg import sql +from generalresearch.models.custom_types import UUIDStr from generalresearch.models.thl.user import User if TYPE_CHECKING: - from generalresearch.models.custom_types import UUIDStr + from generalresearch.pg_helper import PostgresConfig logging.basicConfig() diff --git a/generalresearch/managers/thl/user_manager/user_manager.py b/generalresearch/managers/thl/user_manager/user_manager.py index 907a030..52e0567 100644 --- a/generalresearch/managers/thl/user_manager/user_manager.py +++ b/generalresearch/managers/thl/user_manager/user_manager.py @@ -22,12 +22,12 @@ from generalresearch.managers.thl.user_manager.rate_limit import ( from generalresearch.managers.thl.user_manager.redis_user_manager import ( RedisUserManager, ) +from generalresearch.models.custom_types import UUIDStr from generalresearch.utils.copying_cache import deepcopy_return if TYPE_CHECKING: from generalresearch.managers.thl.userhealth import AuditLogManager - from generalresearch.models.custom_types import UUIDStr from generalresearch.models.thl.product import Product from generalresearch.models.thl.user import User from generalresearch.models.thl.userhealth import AuditLog diff --git a/generalresearch/managers/thl/wall.py b/generalresearch/managers/thl/wall.py index 83697f5..bffd4c8 100644 --- a/generalresearch/managers/thl/wall.py +++ b/generalresearch/managers/thl/wall.py @@ -19,6 +19,7 @@ from generalresearch.managers.base import ( PostgresManagerWithRedis, ) from generalresearch.managers.utils import parse_order_by +from generalresearch.models.custom_types import SurveyKey, UUIDStr from generalresearch.models.definitions import Source from generalresearch.models.thl.definitions import ( WallAdjustedStatus, @@ -35,7 +36,6 @@ if TYPE_CHECKING: from generalresearch.managers.base import ( Permission, ) - from generalresearch.models.custom_types import SurveyKey, UUIDStr from generalresearch.models.thl.definitions import ( ReportValue, Status, diff --git a/generalresearch/models/admin/request.py b/generalresearch/models/admin/request.py index f128e1b..6112786 100644 --- a/generalresearch/models/admin/request.py +++ b/generalresearch/models/admin/request.py @@ -2,13 +2,12 @@ from __future__ import annotations from datetime import UTC, datetime, timedelta from enum import Enum -from typing import TYPE_CHECKING, Literal +from typing import Literal import pandas as pd from pydantic import BaseModel, Field, computed_field, model_validator -if TYPE_CHECKING: - from generalresearch.models.custom_types import AwareDatetimeISO +from generalresearch.models.custom_types import AwareDatetimeISO class ReportType(Enum): diff --git a/generalresearch/models/cint/question.py b/generalresearch/models/cint/question.py index ab46653..5f703ee 100644 --- a/generalresearch/models/cint/question.py +++ b/generalresearch/models/cint/question.py @@ -8,6 +8,7 @@ from uuid import UUID from pydantic import BaseModel, Field, field_validator, model_validator +from generalresearch.models.custom_types import AwareDatetimeISO from generalresearch.models.definitions import Source, string_utils from generalresearch.models.thl.profiling.marketplace import ( MarketplaceQuestion, @@ -16,7 +17,6 @@ from generalresearch.models.thl.profiling.marketplace import ( if TYPE_CHECKING: from generalresearch.models.cint import CintQuestionIdType - from generalresearch.models.custom_types import AwareDatetimeISO from generalresearch.models.thl.profiling.upk_question import ( UpkQuestion, ) diff --git a/generalresearch/models/cint/survey.py b/generalresearch/models/cint/survey.py index ebba09e..cd429dd 100644 --- a/generalresearch/models/cint/survey.py +++ b/generalresearch/models/cint/survey.py @@ -18,6 +18,11 @@ from pydantic import ( ) from generalresearch.locales import Localelator +from generalresearch.models.custom_types import ( + AlphaNumStr, + AwareDatetimeISO, + CoercedStr, +) from generalresearch.models.definitions import Source, TaskCalculationType from generalresearch.models.thl.demographics import Gender from generalresearch.models.thl.survey import MarketplaceTask @@ -28,11 +33,7 @@ from generalresearch.models.thl.survey.condition import ( if TYPE_CHECKING: from generalresearch.models.cint import CintQuestionIdType - from generalresearch.models.custom_types import ( - AlphaNumStr, - AwareDatetimeISO, - CoercedStr, - ) + logging.basicConfig() logger = logging.getLogger() diff --git a/generalresearch/models/dynata/question.py b/generalresearch/models/dynata/question.py index 60c7366..cbd5d86 100644 --- a/generalresearch/models/dynata/question.py +++ b/generalresearch/models/dynata/question.py @@ -7,16 +7,14 @@ import re from datetime import timedelta from enum import StrEnum from functools import cached_property -from typing import TYPE_CHECKING, Any, Literal +from typing import Any, Literal from pydantic import BaseModel, Field, PositiveInt, field_validator, model_validator +from generalresearch.models.custom_types import AwareDatetimeISO from generalresearch.models.definitions import MAX_INT32, Source from generalresearch.models.thl.profiling.marketplace import MarketplaceQuestion -if TYPE_CHECKING: - from generalresearch.models.custom_types import AwareDatetimeISO - logging.basicConfig() logger = logging.getLogger() logger.setLevel(logging.INFO) diff --git a/generalresearch/models/dynata/survey.py b/generalresearch/models/dynata/survey.py index 4174d31..491157b 100644 --- a/generalresearch/models/dynata/survey.py +++ b/generalresearch/models/dynata/survey.py @@ -19,6 +19,13 @@ from pydantic import ( ) from generalresearch.locales import Localelator +from generalresearch.models.custom_types import ( + AlphaNumStr, + AlphaNumStrSet, + AwareDatetimeISO, + CoercedStr, + DeviceTypes, +) from generalresearch.models.definitions import Source from generalresearch.models.dynata import DynataStatus from generalresearch.models.thl.demographics import ( @@ -31,13 +38,7 @@ from generalresearch.models.thl.survey.condition import ( ) if TYPE_CHECKING: - from generalresearch.models.custom_types import ( - AlphaNumStr, - AlphaNumStrSet, - AwareDatetimeISO, - CoercedStr, - DeviceTypes, - ) + from generalresearch.models.definitions import TaskCalculationType logging.basicConfig() diff --git a/generalresearch/models/events.py b/generalresearch/models/events.py index 34f6be8..7014d22 100644 --- a/generalresearch/models/events.py +++ b/generalresearch/models/events.py @@ -13,12 +13,14 @@ from pydantic import ( model_validator, ) +from generalresearch.models.custom_types import ( + AwareDatetimeISO, + CountryISOLike, + UUIDStr, +) + if TYPE_CHECKING: - from generalresearch.models.custom_types import ( - AwareDatetimeISO, - CountryISOLike, - UUIDStr, - ) + from generalresearch.models.definitions import Source from generalresearch.models.thl.definitions import ( SessionStatusCode2, diff --git a/generalresearch/models/gr/__init__.py b/generalresearch/models/gr/__init__.py index 79f05d3..7e1516b 100644 --- a/generalresearch/models/gr/__init__.py +++ b/generalresearch/models/gr/__init__.py @@ -1,13 +1,13 @@ -# from generalresearch.models.gr.authentication import GRToken, GRUser -# from generalresearch.models.gr.business import Business -# from generalresearch.models.gr.team import Team -# from generalresearch.models.thl.finance import BusinessBalances -# from generalresearch.models.thl.payout import BrokerageProductPayoutEvent -# from generalresearch.models.thl.product import Product +from generalresearch.models.gr.authentication import GRToken, GRUser +from generalresearch.models.gr.business import Business +from generalresearch.models.gr.team import Team +from generalresearch.models.thl.finance import BusinessBalances +from generalresearch.models.thl.payout import BrokerageProductPayoutEvent +from generalresearch.models.thl.product import Product -# _ = Business, Product, BrokerageProductPayoutEvent, BusinessBalances +_ = Business, Product, BrokerageProductPayoutEvent, BusinessBalances -# GRUser.model_rebuild() -# GRToken.model_rebuild() -# Business.model_rebuild() -# Team.model_rebuild() +GRUser.model_rebuild() +GRToken.model_rebuild() +Business.model_rebuild() +Team.model_rebuild() diff --git a/generalresearch/models/gr/authentication.py b/generalresearch/models/gr/authentication.py index 21ece8c..41c3eaa 100644 --- a/generalresearch/models/gr/authentication.py +++ b/generalresearch/models/gr/authentication.py @@ -17,9 +17,9 @@ from pydantic import ( ) from generalresearch.decorators import LOG +from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr if TYPE_CHECKING: - from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.gr.business import Business from generalresearch.models.gr.team import Team from generalresearch.models.thl.product import Product @@ -164,8 +164,8 @@ class GRUser(BaseModel): self.prefetch_businesses(pg_config=pg_config, redis_config=redis_config) self.prefetch_teams(pg_config=pg_config, redis_config=redis_config) - business_uuids = self.business_uuids - team_uuids = self.team_uuids + business_uuids = self.business_uuids or [] + team_uuids = self.team_uuids or [] if len(business_uuids + team_uuids) == 0: self.products = [] diff --git a/generalresearch/models/gr/business.py b/generalresearch/models/gr/business.py index c6d3468..b01c902 100644 --- a/generalresearch/models/gr/business.py +++ b/generalresearch/models/gr/business.py @@ -31,13 +31,12 @@ from generalresearch.models.custom_types import ( 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.models.thl.ledger import LedgerAccount, OrderBy +from generalresearch.models.thl.payout import BusinessPayoutEvent from generalresearch.utils.aggregation import group_by_year if TYPE_CHECKING: from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge - from generalresearch.models.thl.ledger import LedgerAccount - from generalresearch.models.thl.payout import BusinessPayoutEvent from generalresearch.pg_helper import PostgresConfig from generalresearch.redis_helper import RedisConfig diff --git a/generalresearch/models/gr/team.py b/generalresearch/models/gr/team.py index b36ac4c..b1553e2 100644 --- a/generalresearch/models/gr/team.py +++ b/generalresearch/models/gr/team.py @@ -23,6 +23,11 @@ from pydantic.json_schema import SkipJsonSchema from generalresearch.decorators import LOG from generalresearch.models.admin.request import ReportRequest, ReportType +from generalresearch.models.custom_types import ( + AwareDatetimeISO, + UUIDStr, + UUIDStrCoerce, +) from generalresearch.utils.enum import ReprEnumMeta if TYPE_CHECKING: @@ -39,11 +44,6 @@ if TYPE_CHECKING: from generalresearch.managers.gr.business import BusinessManager from generalresearch.managers.gr.team import MembershipManager from generalresearch.managers.thl.product import ProductManager - from generalresearch.models.custom_types import ( - AwareDatetimeISO, - UUIDStr, - UUIDStrCoerce, - ) from generalresearch.models.gr.authentication import GRUser from generalresearch.models.gr.business import Business from generalresearch.models.thl.product import Product diff --git a/generalresearch/models/innovate/survey.py b/generalresearch/models/innovate/survey.py index e718dda..7232228 100644 --- a/generalresearch/models/innovate/survey.py +++ b/generalresearch/models/innovate/survey.py @@ -24,6 +24,12 @@ from pydantic import ( ) from generalresearch.locales import Localelator +from generalresearch.models.custom_types import ( + AlphaNumStrSet, + AwareDatetimeISO, + CoercedStr, + DeviceTypes, +) from generalresearch.models.definitions import ( LogicalOperator, Source, @@ -41,12 +47,7 @@ from generalresearch.models.thl.survey.condition import ( ) if TYPE_CHECKING: - from generalresearch.models.custom_types import ( - AlphaNumStrSet, - AwareDatetimeISO, - CoercedStr, - DeviceTypes, - ) + from generalresearch.models.definitions import ( TaskCalculationType, ) diff --git a/generalresearch/models/legacy/bucket.py b/generalresearch/models/legacy/bucket.py index 5f53b89..23eeb38 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 TYPE_CHECKING, Any, Literal, Self +from typing import Any, Literal, Self from pydantic import ( BaseModel, @@ -15,16 +15,14 @@ from pydantic import ( model_validator, ) +from generalresearch.models.custom_types import ( + HttpsUrl, + PropertyCode, + UUIDStr, +) from generalresearch.models.definitions import Source from generalresearch.models.thl.stats import StatisticalSummary -if TYPE_CHECKING: - from generalresearch.models.custom_types import ( - HttpsUrl, - PropertyCode, - UUIDStr, - ) - logger = logging.getLogger() Eligibility = Literal["conditional", "unconditional", "ineligible"] diff --git a/generalresearch/models/legacy/offerwall.py b/generalresearch/models/legacy/offerwall.py index c150dcb..013b506 100644 --- a/generalresearch/models/legacy/offerwall.py +++ b/generalresearch/models/legacy/offerwall.py @@ -4,13 +4,14 @@ from typing import TYPE_CHECKING from pydantic import BaseModel, ConfigDict, Field, NonNegativeInt +from generalresearch.models.custom_types import UUIDStr from generalresearch.models.legacy.definitions import OfferwallReason from generalresearch.models.thl.payout_format import ( PayoutFormatField, ) if TYPE_CHECKING: - from generalresearch.models.custom_types import UUIDStr + from generalresearch.models.legacy.bucket import ( BucketBase, MarketplaceBucket, diff --git a/generalresearch/models/legacy/questions.py b/generalresearch/models/legacy/questions.py index c333804..e6803f0 100644 --- a/generalresearch/models/legacy/questions.py +++ b/generalresearch/models/legacy/questions.py @@ -15,6 +15,7 @@ from pydantic import ( ) from sentry_sdk import capture_exception +from generalresearch.models.custom_types import UUIDStr from generalresearch.models.legacy.api_status import StatusResponse if TYPE_CHECKING: @@ -22,7 +23,6 @@ if TYPE_CHECKING: UserManager, ) from generalresearch.managers.thl.wall import WallManager - from generalresearch.models.custom_types import UUIDStr from generalresearch.models.thl.profiling.upk_question import ( UpkQuestionOut, ) diff --git a/generalresearch/models/lucid/survey.py b/generalresearch/models/lucid/survey.py index a04e529..02b31ab 100644 --- a/generalresearch/models/lucid/survey.py +++ b/generalresearch/models/lucid/survey.py @@ -4,6 +4,12 @@ from typing import TYPE_CHECKING, Any, Self from pydantic import BaseModel, ConfigDict, Field, NonNegativeInt +from generalresearch.models.custom_types import ( + AwareDatetimeISO, + BigAutoInteger, + CoercedStr, + UUIDStr, +) from generalresearch.models.definitions import Source from generalresearch.models.thl.survey.condition import ( ConditionValueType, @@ -11,12 +17,6 @@ from generalresearch.models.thl.survey.condition import ( ) if TYPE_CHECKING: - from generalresearch.models.custom_types import ( - AwareDatetimeISO, - BigAutoInteger, - CoercedStr, - UUIDStr, - ) from generalresearch.models.thl.locales import CountryISO, LanguageISO diff --git a/generalresearch/models/morning/survey.py b/generalresearch/models/morning/survey.py index 25accb6..01cc4ff 100644 --- a/generalresearch/models/morning/survey.py +++ b/generalresearch/models/morning/survey.py @@ -25,6 +25,10 @@ from pydantic import ( ) from generalresearch.locales import Localelator +from generalresearch.models.custom_types import ( + AwareDatetimeISO, + UUIDStrCoerce, +) from generalresearch.models.definitions import Source from generalresearch.models.morning import MorningStatus from generalresearch.models.thl.demographics import Gender @@ -35,10 +39,7 @@ from generalresearch.models.thl.survey.condition import ( ) if TYPE_CHECKING: - from generalresearch.models.custom_types import ( - AwareDatetimeISO, - UUIDStrCoerce, - ) + from generalresearch.models.morning import MorningQuestionID from generalresearch.models.morning.question import MorningQuestion from generalresearch.models.thl.locales import ( diff --git a/generalresearch/models/network/mtr/execute.py b/generalresearch/models/network/mtr/execute.py index 1a7c963..5ab7632 100644 --- a/generalresearch/models/network/mtr/execute.py +++ b/generalresearch/models/network/mtr/execute.py @@ -1,9 +1,9 @@ from __future__ import annotations from datetime import UTC, datetime -from typing import TYPE_CHECKING from uuid import uuid4 +from generalresearch.models.custom_types import UUIDStr from generalresearch.models.network.definitions import IPProtocol from generalresearch.models.network.mtr.command import ( get_mtr_version, @@ -21,9 +21,6 @@ from generalresearch.models.network.tool_run_command import ( ) from generalresearch.models.network.utils import get_source_ip -if TYPE_CHECKING: - from generalresearch.models.custom_types import UUIDStr - def execute_mtr( ip: str, diff --git a/generalresearch/models/network/nmap/execute.py b/generalresearch/models/network/nmap/execute.py index e3610d9..09ec28b 100644 --- a/generalresearch/models/network/nmap/execute.py +++ b/generalresearch/models/network/nmap/execute.py @@ -1,8 +1,8 @@ from __future__ import annotations -from typing import TYPE_CHECKING from uuid import uuid4 +from generalresearch.models.custom_types import UUIDStr from generalresearch.models.network.nmap.command import run_nmap from generalresearch.models.network.tool_run import ( NmapRun, @@ -15,9 +15,6 @@ from generalresearch.models.network.tool_run_command import ( NmapRunCommandOptions, ) -if TYPE_CHECKING: - from generalresearch.models.custom_types import UUIDStr - def execute_nmap( ip: str, diff --git a/generalresearch/models/network/rdns/execute.py b/generalresearch/models/network/rdns/execute.py index 6c14f77..d6de84b 100644 --- a/generalresearch/models/network/rdns/execute.py +++ b/generalresearch/models/network/rdns/execute.py @@ -1,9 +1,9 @@ from __future__ import annotations from datetime import UTC, datetime -from typing import TYPE_CHECKING from uuid import uuid4 +from generalresearch.models.custom_types import UUIDStr from generalresearch.models.network.rdns.command import ( get_dig_version, run_rdns, @@ -19,9 +19,6 @@ from generalresearch.models.network.tool_run_command import ( RDNSRunCommandOptions, ) -if TYPE_CHECKING: - from generalresearch.models.custom_types import UUIDStr - def execute_rdns(ip: str, scan_group_id: UUIDStr | None = None): started_at = datetime.now(tz=UTC) diff --git a/generalresearch/models/network/tool_run.py b/generalresearch/models/network/tool_run.py index 8479f15..9088fe3 100644 --- a/generalresearch/models/network/tool_run.py +++ b/generalresearch/models/network/tool_run.py @@ -6,12 +6,14 @@ from uuid import uuid4 from pydantic import BaseModel, Field, PositiveInt +from generalresearch.models.custom_types import ( + AwareDatetimeISO, + IPvAnyAddressStr, + UUIDStr, +) + if TYPE_CHECKING: - from generalresearch.models.custom_types import ( - AwareDatetimeISO, - IPvAnyAddressStr, - UUIDStr, - ) + from generalresearch.models.network.mtr.result import MTRResult from generalresearch.models.network.nmap.result import NmapResult from generalresearch.models.network.rdns.result import RDNSResult diff --git a/generalresearch/models/precision/survey.py b/generalresearch/models/precision/survey.py index a9e34e6..b77a365 100644 --- a/generalresearch/models/precision/survey.py +++ b/generalresearch/models/precision/survey.py @@ -15,6 +15,13 @@ from pydantic import ( model_validator, ) +from generalresearch.models.custom_types import ( + AlphaNumStrSet, + AwareDatetimeISO, + CoercedStr, + DeviceTypes, + UUIDStrCoerce, +) from generalresearch.models.definitions import Source from generalresearch.models.precision import PrecisionStatus from generalresearch.models.thl.demographics import Gender @@ -25,13 +32,6 @@ from generalresearch.models.thl.survey.condition import ( ) if TYPE_CHECKING: - from generalresearch.models.custom_types import ( - AlphaNumStrSet, - AwareDatetimeISO, - CoercedStr, - DeviceTypes, - UUIDStrCoerce, - ) from generalresearch.models.precision import PrecisionQuestionID diff --git a/generalresearch/models/prodege/survey.py b/generalresearch/models/prodege/survey.py index e3c765e..26898d0 100644 --- a/generalresearch/models/prodege/survey.py +++ b/generalresearch/models/prodege/survey.py @@ -20,6 +20,13 @@ from pydantic import ( ) from generalresearch.locales import Localelator +from generalresearch.models.custom_types import ( + AlphaNumStrSet, + AwareDatetimeISO, + CoercedStr, + InclExcl, + UUIDStr, +) from generalresearch.models.definitions import ( LogicalOperator, Source, @@ -38,13 +45,7 @@ from generalresearch.models.thl.survey.condition import ( ) if TYPE_CHECKING: - from generalresearch.models.custom_types import ( - AlphaNumStrSet, - AwareDatetimeISO, - CoercedStr, - InclExcl, - UUIDStr, - ) + from generalresearch.models.prodege import ( ProdegeQuestionIdType, ProdgeRedirectStatus, diff --git a/generalresearch/models/repdata/question.py b/generalresearch/models/repdata/question.py index a578741..0ec102b 100644 --- a/generalresearch/models/repdata/question.py +++ b/generalresearch/models/repdata/question.py @@ -17,11 +17,11 @@ from pydantic import ( model_validator, ) +from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.definitions import MAX_INT32, Source from generalresearch.models.thl.profiling.marketplace import MarketplaceQuestion if TYPE_CHECKING: - from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.thl.profiling.upk_question import ( UpkQuestion, ) diff --git a/generalresearch/models/repdata/survey.py b/generalresearch/models/repdata/survey.py index fc1b649..a69a54b 100644 --- a/generalresearch/models/repdata/survey.py +++ b/generalresearch/models/repdata/survey.py @@ -6,7 +6,7 @@ import logging from datetime import UTC, datetime from decimal import Decimal from functools import cached_property -from typing import TYPE_CHECKING, Any, Literal, Self +from typing import Any, Literal, Self from uuid import UUID from pydantic import ( @@ -21,6 +21,11 @@ from pydantic import ( from generalresearch.grpc import timestamp_from_datetime from generalresearch.locales import Localelator +from generalresearch.models.custom_types import ( + AwareDatetimeISO, + CoercedStr, + UUIDStr, +) from generalresearch.models.definitions import ( DeviceType, LogicalOperator, @@ -35,13 +40,6 @@ from generalresearch.models.thl.survey.condition import ( MarketplaceCondition, ) -if TYPE_CHECKING: - from generalresearch.models.custom_types import ( - AwareDatetimeISO, - CoercedStr, - UUIDStr, - ) - logging.basicConfig() logger = logging.getLogger() logger.setLevel(logging.INFO) diff --git a/generalresearch/models/spectrum/survey.py b/generalresearch/models/spectrum/survey.py index a02c510..6689d27 100644 --- a/generalresearch/models/spectrum/survey.py +++ b/generalresearch/models/spectrum/survey.py @@ -4,12 +4,19 @@ import json import logging from datetime import UTC from decimal import Decimal -from typing import TYPE_CHECKING, Any, Literal, Self +from typing import Any, Literal, Self from more_itertools import flatten from pydantic import BaseModel, ConfigDict, Field, computed_field, model_validator from generalresearch.locales import Localelator +from generalresearch.models.custom_types import ( + AlphaNumStr, + AlphaNumStrSet, + AwareDatetimeISO, + CoercedStr, + UUIDStrSet, +) from generalresearch.models.definitions import Source, TaskCalculationType from generalresearch.models.spectrum import SpectrumStatus from generalresearch.models.thl.demographics import Gender @@ -19,15 +26,6 @@ from generalresearch.models.thl.survey.condition import ( MarketplaceCondition, ) -if TYPE_CHECKING: - from generalresearch.models.custom_types import ( - AlphaNumStr, - AlphaNumStrSet, - AwareDatetimeISO, - CoercedStr, - UUIDStrSet, - ) - logging.basicConfig() logger = logging.getLogger() logger.setLevel(logging.INFO) diff --git a/generalresearch/models/thl/__init__.py b/generalresearch/models/thl/__init__.py index 45278f8..d0791e4 100644 --- a/generalresearch/models/thl/__init__.py +++ b/generalresearch/models/thl/__init__.py @@ -1,21 +1,23 @@ -# from generalresearch.models.thl.finance import ( -# POPFinancial, -# ProductBalances, -# ) -# from generalresearch.models.thl.payout import ( -# # BrokerageProductPayoutEvent, -# PayoutEvent, -# ) -# from generalresearch.models.thl.product import Product +from generalresearch.models.thl.finance import ( + POPFinancial, + ProductBalances, +) +from generalresearch.models.thl.ledger import LedgerAccount +from generalresearch.models.thl.payout import ( + BrokerageProductPayoutEvent, + PayoutEvent, +) +from generalresearch.models.thl.product import Product -# _ = ( -# Product, -# PayoutEvent, -# BrokerageProductPayoutEvent, -# ProductBalances, -# POPFinancial, -# ) +_ = ( + Product, + PayoutEvent, + BrokerageProductPayoutEvent, + ProductBalances, + POPFinancial, +) -# Product.model_rebuild() -# PayoutEvent.model_rebuild() -# BrokerageProductPayoutEvent.model_rebuild() +Product.model_rebuild() +LedgerAccount.model_rebuild() +PayoutEvent.model_rebuild() +BrokerageProductPayoutEvent.model_rebuild() diff --git a/generalresearch/models/thl/contest/__init__.py b/generalresearch/models/thl/contest/__init__.py index f243ce4..c8342b3 100644 --- a/generalresearch/models/thl/contest/__init__.py +++ b/generalresearch/models/thl/contest/__init__.py @@ -12,11 +12,11 @@ from pydantic import ( model_validator, ) +from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.thl.contest.definitions import ContestPrizeKind if TYPE_CHECKING: from generalresearch.currency import USDCent - from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.thl.user import User diff --git a/generalresearch/models/thl/contest/contest.py b/generalresearch/models/thl/contest/contest.py index 173d486..5e30778 100644 --- a/generalresearch/models/thl/contest/contest.py +++ b/generalresearch/models/thl/contest/contest.py @@ -15,6 +15,7 @@ from pydantic import ( model_validator, ) +from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.thl.contest import ( ContestEndCondition, ContestPrize, @@ -26,7 +27,6 @@ from generalresearch.models.thl.contest.definitions import ( ) if TYPE_CHECKING: - from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.thl.contest import ( ContestWinner, ) diff --git a/generalresearch/models/thl/contest/contest_entry.py b/generalresearch/models/thl/contest/contest_entry.py index 17b288b..261b3fc 100644 --- a/generalresearch/models/thl/contest/contest_entry.py +++ b/generalresearch/models/thl/contest/contest_entry.py @@ -12,12 +12,12 @@ from pydantic import ( ) from generalresearch.currency import USDCent +from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.thl.contest.definitions import ( ContestEntryType, ) if TYPE_CHECKING: - from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.thl.user import User diff --git a/generalresearch/models/thl/contest/milestone.py b/generalresearch/models/thl/contest/milestone.py index e2ff2bc..db5ba2f 100644 --- a/generalresearch/models/thl/contest/milestone.py +++ b/generalresearch/models/thl/contest/milestone.py @@ -2,7 +2,7 @@ from __future__ import annotations import logging from datetime import timedelta -from typing import TYPE_CHECKING, Any, Literal, Self +from typing import Any, Literal, Self from pydantic import ( BaseModel, @@ -13,6 +13,7 @@ from pydantic import ( ) from generalresearch.currency import USDCent +from generalresearch.models.custom_types import AwareDatetimeISO from generalresearch.models.thl.contest import ( ContestPrize, ) @@ -31,9 +32,6 @@ from generalresearch.models.thl.contest.definitions import ( ContestType, ) -if TYPE_CHECKING: - from generalresearch.models.custom_types import AwareDatetimeISO - logging.basicConfig() LOG = logging.getLogger() LOG.setLevel(logging.INFO) diff --git a/generalresearch/models/thl/finance.py b/generalresearch/models/thl/finance.py index 9e7d2c3..f012cbf 100644 --- a/generalresearch/models/thl/finance.py +++ b/generalresearch/models/thl/finance.py @@ -18,6 +18,7 @@ from pydantic import ( from pydantic.json_schema import SkipJsonSchema from generalresearch.config import is_debug +from generalresearch.currency import USDCent from generalresearch.decorators import LOG from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.thl.definitions import SessionAdjustedStatus @@ -26,7 +27,7 @@ payout_example = random.randint(150, 750 * 100) 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.models.thl.product import Product @@ -324,8 +325,6 @@ class ProductBalances(BaseModel): ) @property def payout_usd_str(self) -> str: - from generalresearch.currency import USDCent - return USDCent(self.payout).to_usd_str() @computed_field( @@ -384,8 +383,6 @@ class ProductBalances(BaseModel): ) @property def payment_usd_str(self): - from generalresearch.currency import USDCent - return USDCent(self.payment).to_usd_str() @computed_field( @@ -427,8 +424,6 @@ class ProductBalances(BaseModel): ) @property def retainer_usd_str(self) -> str: - from generalresearch.currency import USDCent - return USDCent(self.retainer).to_usd_str() @computed_field( @@ -460,8 +455,6 @@ class ProductBalances(BaseModel): ) @property def available_balance_usd_str(self) -> str: - from generalresearch.currency import USDCent - return USDCent(self.available_balance).to_usd_str() @computed_field( @@ -477,8 +470,6 @@ class ProductBalances(BaseModel): ) @property def recoup(self) -> USDCent: - from generalresearch.currency import USDCent - if self.balance >= 0: return USDCent(0) @@ -578,8 +569,6 @@ class BusinessBalances(BaseModel): ) @property def payout_usd_str(self) -> str: - from generalresearch.currency import USDCent - return USDCent(self.payout).to_usd_str() @computed_field( @@ -681,8 +670,6 @@ class BusinessBalances(BaseModel): ) @property def payment_usd_str(self) -> str: - from generalresearch.currency import USDCent - return USDCent(self.payment).to_usd_str() @computed_field( @@ -730,8 +717,6 @@ class BusinessBalances(BaseModel): ) @property def retainer_usd_str(self) -> str: - from generalresearch.currency import USDCent - return USDCent(self.retainer).to_usd_str() @computed_field( @@ -763,8 +748,6 @@ class BusinessBalances(BaseModel): ) @property def available_balance_usd_str(self) -> str: - from generalresearch.currency import USDCent - return USDCent(self.available_balance).to_usd_str() # --- Properties: account related --- @@ -799,8 +782,6 @@ class BusinessBalances(BaseModel): """Returns the sum of this Business' recouped amount from any children Products. """ - from generalresearch.currency import USDCent - return USDCent(sum([i.recoup for i in self.product_balances])) @computed_field( diff --git a/generalresearch/models/thl/ledger.py b/generalresearch/models/thl/ledger.py index fbfb6bb..af9d042 100644 --- a/generalresearch/models/thl/ledger.py +++ b/generalresearch/models/thl/ledger.py @@ -16,7 +16,12 @@ from pydantic import ( model_validator, ) -from generalresearch.models.custom_types import check_valid_uuid +from generalresearch.models.custom_types import ( + AwareDatetimeISO, + HttpsUrlStr, + UUIDStr, + check_valid_uuid, +) from generalresearch.models.thl.pagination import Page from generalresearch.models.thl.payout_format import ( PayoutFormatType, @@ -25,11 +30,6 @@ from generalresearch.models.thl.payout_format import ( from generalresearch.utils.enum import ReprEnumMeta if TYPE_CHECKING: - from generalresearch.models.custom_types import ( - AwareDatetimeISO, - HttpsUrlStr, - UUIDStr, - ) from generalresearch.models.thl.payout_format import ( PayoutFormatType, ) diff --git a/generalresearch/models/thl/offerwall/base.py b/generalresearch/models/thl/offerwall/base.py index fb0bc77..c99df16 100644 --- a/generalresearch/models/thl/offerwall/base.py +++ b/generalresearch/models/thl/offerwall/base.py @@ -19,6 +19,7 @@ from pydantic import ( model_validator, ) +from generalresearch.models.custom_types import HttpsUrl, UUIDStr from generalresearch.models.definitions import Source from generalresearch.models.legacy.bucket import ( Bucket as LegacyBucket, @@ -38,7 +39,6 @@ from generalresearch.models.thl.offerwall.bucket import ( from generalresearch.models.thl.soft_pair import SoftPairResultType if TYPE_CHECKING: - from generalresearch.models.custom_types import HttpsUrl, UUIDStr from generalresearch.models.legacy.bucket import ( CategoryAssociation, Eligibility, diff --git a/generalresearch/models/thl/offerwall/cache.py b/generalresearch/models/thl/offerwall/cache.py index 2a733c9..9b2f462 100644 --- a/generalresearch/models/thl/offerwall/cache.py +++ b/generalresearch/models/thl/offerwall/cache.py @@ -5,8 +5,9 @@ from typing import TYPE_CHECKING, Any from pydantic import BaseModel, Field +from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr + if TYPE_CHECKING: - 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 ( diff --git a/generalresearch/models/thl/payout.py b/generalresearch/models/thl/payout.py index 128723b..8759cd4 100644 --- a/generalresearch/models/thl/payout.py +++ b/generalresearch/models/thl/payout.py @@ -17,19 +17,17 @@ from pydantic import ( from pydantic.json_schema import SkipJsonSchema from generalresearch.currency import USDCent +from generalresearch.models.custom_types import ( + AwareDatetimeISO, + UUIDStr, + UUIDStrCoerce, +) from generalresearch.models.thl.definitions import PayoutStatus +from generalresearch.models.thl.wallet.cashout_method import ( + CashMailOrderData, +) from generalresearch.models.thl.wallet.definitions import PayoutType -if TYPE_CHECKING: - from generalresearch.models.custom_types import ( - AwareDatetimeISO, - UUIDStr, - UUIDStrCoerce, - ) - from generalresearch.models.thl.wallet.cashout_method import ( - CashMailOrderData, - ) - class PayoutEvent(BaseModel): """Base Pydantic Model to represent the `event_payout` table diff --git a/generalresearch/models/thl/profiling/marketplace.py b/generalresearch/models/thl/profiling/marketplace.py index 23501e3..19d45d6 100644 --- a/generalresearch/models/thl/profiling/marketplace.py +++ b/generalresearch/models/thl/profiling/marketplace.py @@ -7,15 +7,16 @@ from typing import TYPE_CHECKING, Any from pydantic import BaseModel, ConfigDict, Field, PositiveInt, computed_field +from generalresearch.models.custom_types import ( + AwareDatetimeISO, + CountryISOLike, + LanguageISOLike, + UUIDStr, +) from generalresearch.models.definitions import MAX_INT32 if TYPE_CHECKING: - 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/question.py b/generalresearch/models/thl/profiling/question.py index 920dd3a..ea6a3c9 100644 --- a/generalresearch/models/thl/profiling/question.py +++ b/generalresearch/models/thl/profiling/question.py @@ -9,13 +9,14 @@ from pydantic import ( computed_field, ) +from generalresearch.models.custom_types import ( + AwareDatetimeISO, + CountryISOLike, + LanguageISOLike, + UUIDStr, +) + if TYPE_CHECKING: - from generalresearch.models.custom_types import ( - AwareDatetimeISO, - CountryISOLike, - LanguageISOLike, - UUIDStr, - ) from generalresearch.models.thl.profiling.upk_question import UpkQuestion diff --git a/generalresearch/models/thl/profiling/upk_property.py b/generalresearch/models/thl/profiling/upk_property.py index 922f5a4..96f1b4c 100644 --- a/generalresearch/models/thl/profiling/upk_property.py +++ b/generalresearch/models/thl/profiling/upk_property.py @@ -7,10 +7,10 @@ from uuid import uuid4 from pydantic import BaseModel, ConfigDict, Field, TypeAdapter +from generalresearch.models.custom_types import CountryISOLike, UUIDStr from generalresearch.utils.enum import ReprEnumMeta if TYPE_CHECKING: - from generalresearch.models.custom_types import CountryISOLike, UUIDStr from generalresearch.models.thl.category import Category diff --git a/generalresearch/models/thl/profiling/upk_question_answer.py b/generalresearch/models/thl/profiling/upk_question_answer.py index 4d07970..f25baa7 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 UTC, datetime -from typing import TYPE_CHECKING, Any, Self +from typing import Any, Self from uuid import uuid4 from pydantic import ( @@ -13,19 +13,17 @@ from pydantic import ( model_validator, ) +from generalresearch.models.custom_types import ( + AwareDatetimeISO, + CountryISOLike, + UUIDStr, +) from generalresearch.models.definitions import MAX_INT32 from generalresearch.models.thl.profiling.upk_property import ( Cardinality, PropertyType, ) -if TYPE_CHECKING: - from generalresearch.models.custom_types import ( - AwareDatetimeISO, - CountryISOLike, - UUIDStr, - ) - class UpkQuestionAnswer(BaseModel): diff --git a/generalresearch/models/thl/profiling/user_info.py b/generalresearch/models/thl/profiling/user_info.py index 40b4b17..5124d17 100644 --- a/generalresearch/models/thl/profiling/user_info.py +++ b/generalresearch/models/thl/profiling/user_info.py @@ -5,8 +5,9 @@ from typing import TYPE_CHECKING from pydantic import BaseModel, ConfigDict, Field from pydantic.json_schema import SkipJsonSchema +from generalresearch.models.custom_types import AwareDatetimeISO + if TYPE_CHECKING: - 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/session.py b/generalresearch/models/thl/session.py index 65b885e..5812c75 100644 --- a/generalresearch/models/thl/session.py +++ b/generalresearch/models/thl/session.py @@ -18,6 +18,12 @@ from pydantic import ( model_validator, ) +from generalresearch.models.custom_types import ( + AwareDatetimeISO, + EnumNameSerializer, + IPvAnyAddressStr, + UUIDStr, +) from generalresearch.models.definitions import Source from generalresearch.models.thl.definitions import ( WALL_ALLOWED_STATUS_CODE_1_2, @@ -37,12 +43,6 @@ if TYPE_CHECKING: from generalresearch.managers.thl.ledger_manager.thl_ledger import ( ThlLedgerManager, ) - 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 ( diff --git a/generalresearch/models/thl/survey/buyer.py b/generalresearch/models/thl/survey/buyer.py index ef309d1..6a67ed7 100644 --- a/generalresearch/models/thl/survey/buyer.py +++ b/generalresearch/models/thl/survey/buyer.py @@ -3,7 +3,7 @@ from __future__ import annotations from datetime import UTC, datetime from decimal import Decimal from math import log -from typing import TYPE_CHECKING, Annotated +from typing import Annotated from pydantic import ( BaseModel, @@ -16,15 +16,13 @@ from pydantic import ( ) from scipy.stats import beta as beta_dist +from generalresearch.models.custom_types import ( + AwareDatetimeISO, + CountryISOLike, + UUIDStr, +) from generalresearch.models.definitions import Source -if TYPE_CHECKING: - from generalresearch.models.custom_types import ( - AwareDatetimeISO, - CountryISOLike, - UUIDStr, - ) - class Buyer(BaseModel): """ diff --git a/generalresearch/models/thl/survey/model.py b/generalresearch/models/thl/survey/model.py index 8986e4d..f8e5083 100644 --- a/generalresearch/models/thl/survey/model.py +++ b/generalresearch/models/thl/survey/model.py @@ -17,17 +17,18 @@ from pydantic import ( ) from generalresearch.managers.thl.buyer import Buyer +from generalresearch.models.custom_types import ( + AwareDatetimeISO, + CountryISOLike, + EnumNameSerializer, + PropertyCode, + SurveyKey, +) from generalresearch.models.thl.definitions import StatusCode1 from generalresearch.models.thl.pagination import Page if TYPE_CHECKING: - from generalresearch.models.custom_types import ( - AwareDatetimeISO, - CountryISOLike, - EnumNameSerializer, - 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 25e07cf..e989bff 100644 --- a/generalresearch/models/thl/survey/penalty.py +++ b/generalresearch/models/thl/survey/penalty.py @@ -6,11 +6,13 @@ from typing import TYPE_CHECKING, Annotated, Literal from pydantic import BaseModel, ConfigDict, Field, TypeAdapter +from generalresearch.models.custom_types import ( + AwareDatetimeISO, + UUIDStr, +) + if TYPE_CHECKING: - from generalresearch.models.custom_types import ( - AwareDatetimeISO, - UUIDStr, - ) + from generalresearch.models.definitions import Source diff --git a/generalresearch/models/thl/task_adjustment.py b/generalresearch/models/thl/task_adjustment.py index fee2007..6e5a935 100644 --- a/generalresearch/models/thl/task_adjustment.py +++ b/generalresearch/models/thl/task_adjustment.py @@ -7,13 +7,13 @@ from uuid import uuid4 from pydantic import BaseModel, ConfigDict, Field, PositiveInt, model_validator +from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.definitions import MAX_INT32 from generalresearch.models.thl.definitions import ( WallAdjustedStatus, ) if TYPE_CHECKING: - from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.definitions import Source diff --git a/generalresearch/models/thl/task_status.py b/generalresearch/models/thl/task_status.py index 6cff884..16ad66f 100644 --- a/generalresearch/models/thl/task_status.py +++ b/generalresearch/models/thl/task_status.py @@ -13,6 +13,11 @@ from pydantic import ( model_validator, ) +from generalresearch.models.custom_types import ( + AwareDatetimeISO, + EnumNameSerializer, + UUIDStr, +) from generalresearch.models.thl.definitions import ( SessionAdjustedStatus, SessionStatusCode2, @@ -27,11 +32,7 @@ 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 ( - AwareDatetimeISO, - EnumNameSerializer, - UUIDStr, - ) + from generalresearch.models.thl.payout_format import ( PayoutFormatType, ) diff --git a/generalresearch/models/thl/user.py b/generalresearch/models/thl/user.py index 1f88dc6..4f94270 100644 --- a/generalresearch/models/thl/user.py +++ b/generalresearch/models/thl/user.py @@ -20,6 +20,7 @@ from pydantic import ( ) from sentry_sdk import set_tag, set_user +from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.definitions import MAX_INT32 if TYPE_CHECKING: @@ -27,7 +28,6 @@ if TYPE_CHECKING: ThlLedgerManager, ) from generalresearch.managers.thl.userhealth import AuditLogManager - from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.thl.ipinfo import GeoIPInformation from generalresearch.models.thl.ledger import LedgerTransaction from generalresearch.models.thl.product import Product diff --git a/generalresearch/models/thl/user_profile.py b/generalresearch/models/thl/user_profile.py index c47c6f2..11dfca1 100644 --- a/generalresearch/models/thl/user_profile.py +++ b/generalresearch/models/thl/user_profile.py @@ -13,10 +13,11 @@ from pydantic import ( ) from pydantic.json_schema import SkipJsonSchema +from generalresearch.models.custom_types import UUIDStr from generalresearch.models.definitions import MAX_INT32, Source if TYPE_CHECKING: - from generalresearch.models.custom_types import UUIDStr + from generalresearch.models.thl.user import User from generalresearch.models.thl.user_streak import UserStreak diff --git a/generalresearch/models/thl/user_quality_event.py b/generalresearch/models/thl/user_quality_event.py index 8c2e25f..4d5db9d 100644 --- a/generalresearch/models/thl/user_quality_event.py +++ b/generalresearch/models/thl/user_quality_event.py @@ -7,11 +7,11 @@ from typing import TYPE_CHECKING, Literal from pydantic import BaseModel, Field, PositiveInt +from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.definitions import MAX_INT32, Source from generalresearch.utils.enum import ReprEnumMeta if TYPE_CHECKING: - from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.thl.definitions import WallAdjustedStatus from generalresearch.models.thl.user import BPUIDStr diff --git a/generalresearch/models/thl/wallet/cashout_method.py b/generalresearch/models/thl/wallet/cashout_method.py index 9383c36..286a126 100644 --- a/generalresearch/models/thl/wallet/cashout_method.py +++ b/generalresearch/models/thl/wallet/cashout_method.py @@ -17,18 +17,19 @@ from pydantic import ( model_validator, ) +from generalresearch.currency import USDCent +from generalresearch.models.custom_types import ( + AwareDatetimeISO, + HttpsUrlStr, + UUIDStr, +) from generalresearch.models.legacy.api_status import StatusResponse from generalresearch.models.thl.definitions import PayoutStatus from generalresearch.models.thl.wallet.definitions import PayoutType from generalresearch.utils.enum import ReprEnumMeta if TYPE_CHECKING: - from generalresearch.currency import USDCent - from generalresearch.models.custom_types import ( - AwareDatetimeISO, - HttpsUrlStr, - UUIDStr, - ) + from generalresearch.models.thl.locales import CountryISO from generalresearch.models.thl.user import BPUIDStr, User from generalresearch.models.thl.wallet.definitions import Currency diff --git a/generalresearch/models/thl/wallet/payout.py b/generalresearch/models/thl/wallet/payout.py index 79c50e1..1fc0f77 100644 --- a/generalresearch/models/thl/wallet/payout.py +++ b/generalresearch/models/thl/wallet/payout.py @@ -3,7 +3,7 @@ from __future__ import annotations import json from collections.abc import Collection from datetime import UTC, datetime -from typing import TYPE_CHECKING, Any +from typing import Any from uuid import uuid4 from pydantic import ( @@ -15,15 +15,13 @@ from pydantic import ( ) from generalresearch.currency import USDCent +from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.thl.definitions import PayoutStatus +from generalresearch.models.thl.wallet.cashout_method import ( + CashMailOrderData, +) from generalresearch.models.thl.wallet.definitions import PayoutType -if TYPE_CHECKING: - from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr - from generalresearch.models.thl.wallet.cashout_method import ( - CashMailOrderData, - ) - class PayoutEvent(BaseModel, validate_assignment=True): """A user has requested to be paid from their wallet balance.""" diff --git a/test_utils/models/gr/conftest.py b/test_utils/models/gr/conftest.py index e493f20..b87f3bb 100644 --- a/test_utils/models/gr/conftest.py +++ b/test_utils/models/gr/conftest.py @@ -9,6 +9,8 @@ import pytest from pydantic import PositiveInt from pydantic_extra_types.phone_numbers import PhoneNumber +from generalresearch.models.custom_types import UUIDStr + if TYPE_CHECKING: from generalresearch.managers.gr.authentication import GRTokenManager, GRUserManager from generalresearch.managers.gr.business import ( @@ -17,7 +19,6 @@ if TYPE_CHECKING: 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, diff --git a/test_utils/models/thl/conftest.py b/test_utils/models/thl/conftest.py index 3c77e27..3545509 100644 --- a/test_utils/models/thl/conftest.py +++ b/test_utils/models/thl/conftest.py @@ -13,6 +13,11 @@ import pytest from grip_client.enums import AccessType from pydantic import PositiveInt +from generalresearch.models.custom_types import ( + AwareDatetimeISO, + IPvAnyAddressStr, + UUIDStr, +) from generalresearch.models.thl.definitions import PayoutStatus from generalresearch.models.thl.session import ( Source, @@ -33,11 +38,6 @@ if TYPE_CHECKING: from generalresearch.managers.thl.user_manager.user_manager import UserManager from generalresearch.managers.thl.userhealth import AuditLogManager, IPRecordManager from generalresearch.managers.thl.wall import WallManager - from generalresearch.models.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 diff --git a/tests/managers/thl/test_ledger/test_lm_accounts.py b/tests/managers/thl/test_ledger/test_lm_accounts.py index cdef99a..57a2261 100644 --- a/tests/managers/thl/test_ledger/test_lm_accounts.py +++ b/tests/managers/thl/test_ledger/test_lm_accounts.py @@ -14,14 +14,14 @@ 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 +from generalresearch.models.custom_types import AccountType, Direction, UUIDStr from generalresearch.models.thl.ledger import ( LedgerAccount, LedgerEntry, ) if TYPE_CHECKING: - from generalresearch.models.custom_types import UUIDStr + from generalresearch.models.thl.ledger import ( LedgerTransaction, ) diff --git a/tests/models/custom_types/test_aware_datetime.py b/tests/models/custom_types/test_aware_datetime.py index 54a5d9b..e8a5aa3 100644 --- a/tests/models/custom_types/test_aware_datetime.py +++ b/tests/models/custom_types/test_aware_datetime.py @@ -2,14 +2,12 @@ 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 -if TYPE_CHECKING: - from generalresearch.models.custom_types import AwareDatetimeISO +from generalresearch.models.custom_types import AwareDatetimeISO logger = logging.getLogger() diff --git a/tests/models/custom_types/test_uuid_str.py b/tests/models/custom_types/test_uuid_str.py index 92489a0..02e6a8b 100644 --- a/tests/models/custom_types/test_uuid_str.py +++ b/tests/models/custom_types/test_uuid_str.py @@ -1,13 +1,11 @@ from __future__ import annotations -from typing import TYPE_CHECKING from uuid import uuid4 import pytest from pydantic import BaseModel, Field, ValidationError -if TYPE_CHECKING: - from generalresearch.models.custom_types import UUIDStr +from generalresearch.models.custom_types import UUIDStr class UUIDStrModel(BaseModel): diff --git a/tests/models/thl/test_product.py b/tests/models/thl/test_product.py index a1b3688..446b59f 100644 --- a/tests/models/thl/test_product.py +++ b/tests/models/thl/test_product.py @@ -631,8 +631,6 @@ class TestProductFinancials: create_main_accounts() delete_df_collection(coll=ledger_collection) - from generalresearch.currency import USDCent - p1: Product = product_factory(business=gr_business) u1: User = user_factory(product=p1) bp_wallet = thl_ledger_manager.get_account_or_create_bp_wallet(product=p1) @@ -714,6 +712,8 @@ class TestProductFinancials: # -- Now pay them out... + from generalresearch.currency import USDCent + bp_payout_factory( product=p1, amount=USDCent(50), -- cgit v1.2.3 From 4e9e08718884b1c4394d16055ba3f30c790ef8d0 Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Tue, 1 Sep 2026 16:58:01 -0700 Subject: Latest state for Django Migrations / Discussion --- generalresearch/grliq/managers/forensic_results.py | 6 +-- generalresearch/models/gr/business.py | 4 +- .../0010_supplierpayout_payout_supplier_payout.py | 51 ++++++++++++++++++++ test_utils/conftest.py | 51 ++++++++++++-------- test_utils/managers/gr/conftest.py | 2 + test_utils/managers/thl/conftest.py | 34 +++++++------ test_utils/models/gr/conftest.py | 55 ++++++++++++++-------- tests/managers/gr/test_business.py | 10 ++-- tests/models/gr/test_business.py | 23 +++++---- tests/test_postgres.py | 13 ++++- 10 files changed, 171 insertions(+), 78 deletions(-) create mode 100644 generalresearch/thl_django/migrations/0010_supplierpayout_payout_supplier_payout.py (limited to 'test_utils/models') diff --git a/generalresearch/grliq/managers/forensic_results.py b/generalresearch/grliq/managers/forensic_results.py index 587b768..158e582 100644 --- a/generalresearch/grliq/managers/forensic_results.py +++ b/generalresearch/grliq/managers/forensic_results.py @@ -1,17 +1,15 @@ from collections.abc import Collection from datetime import datetime -from typing import TYPE_CHECKING, Any +from typing import 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 - class GrlIqCategoryResultsReader: def __init__(self, postgres_config: PostgresConfig): diff --git a/generalresearch/models/gr/business.py b/generalresearch/models/gr/business.py index b01c902..73a2f27 100644 --- a/generalresearch/models/gr/business.py +++ b/generalresearch/models/gr/business.py @@ -449,7 +449,7 @@ class Business(BaseModel): def prebuild_pop_financial( self, - thl_pg_config: PostgresConfig, + product_manager: ProductManager, thl_lm: ThlLedgerManager, ds: GRLDatasets, client: DaskClient, @@ -461,7 +461,7 @@ class Business(BaseModel): financial activity within that time window. """ if self.bp_accounts is None: - self.prefetch_bp_accounts(thl_lm=thl_lm, thl_pg_config=thl_pg_config) + self.prefetch_bp_accounts(thl_lm=thl_lm, product_manager=product_manager) from generalresearch.models.admin.request import ( ReportRequest, diff --git a/generalresearch/thl_django/migrations/0010_supplierpayout_payout_supplier_payout.py b/generalresearch/thl_django/migrations/0010_supplierpayout_payout_supplier_payout.py new file mode 100644 index 0000000..5c3319c --- /dev/null +++ b/generalresearch/thl_django/migrations/0010_supplierpayout_payout_supplier_payout.py @@ -0,0 +1,51 @@ +# Generated by Django 6.1 on 2026-09-01 23:15 + +import django.db.models.deletion +from django.db import migrations, models + + +class Migration(migrations.Migration): + + dependencies = [ + ( + "thl_django", + "0009_toolrun_mtrhop_portscanport_iplabel_mtr_portscan_and_more", + ), + ] + + operations = [ + migrations.CreateModel( + name="SupplierPayout", + fields=[ + ("id", models.BigAutoField(primary_key=True, serialize=False)), + ("ext_ref_id", models.CharField(max_length=64, unique=True)), + ("business_id", models.UUIDField(null=True)), + ("created", models.DateTimeField(auto_now_add=True)), + ("amount", models.BigIntegerField()), + ("status", models.CharField(max_length=20, null=True)), + ("payout_type", models.CharField(max_length=14)), + ("request_data", models.JSONField(null=True)), + ("order_data", models.JSONField(null=True)), + ], + options={ + "db_table": "supplier_payout", + "indexes": [ + models.Index( + fields=["created"], name="supplier_pa_created_336236_idx" + ), + models.Index( + fields=["business_id"], name="supplier_pa_busines_2c7a4e_idx" + ), + ], + }, + ), + migrations.AddField( + model_name="payout", + name="supplier_payout", + field=models.ForeignKey( + null=True, + on_delete=django.db.models.deletion.DO_NOTHING, + to="thl_django.supplierpayout", + ), + ), + ] diff --git a/test_utils/conftest.py b/test_utils/conftest.py index daf6b43..44e36a6 100644 --- a/test_utils/conftest.py +++ b/test_utils/conftest.py @@ -225,8 +225,10 @@ def django_db_factory( _ran = {} import django + from django.apps import apps from django.conf import settings as django_settings from django.core.management import call_command + from django.utils.functional import empty def _inner( django_project: str = "generalresearch.thl_django", @@ -242,34 +244,43 @@ def django_db_factory( # We need model files that are NOT in this repo. gr_path = gr_repo() sys.path.insert(0, str(gr_path)) + print("DJANGO_PROJECT_PATH", str(gr_path), sys.path) # 1. Bootstrapping Django settings - if not django_settings.configured: - django_settings.configure( - DATABASES={ - "default": { - "ENGINE": "django.db.backends.postgresql", - "NAME": postgres_instance_dict["name"], - "USER": postgres_instance_dict["username"], - "PASSWORD": postgres_instance_dict["password"], - "HOST": postgres_instance_dict["host"], - "PORT": postgres_instance_dict["port"], - } - }, - INSTALLED_APPS=[ - "django.contrib.postgres", - "django.contrib.contenttypes", - django_project, - ], - ) + # if not django_settings.configured: + # 1. Reset the lazy wrapper back to an empty state + # if not django_settings.configured: + + django_settings._wrapped = empty + + django_settings.configure( + DATABASES={ + "default": { + "ENGINE": "django.db.backends.postgresql", + "NAME": postgres_instance_dict["name"], + "USER": postgres_instance_dict["username"], + "PASSWORD": postgres_instance_dict["password"], + "HOST": postgres_instance_dict["host"], + "PORT": postgres_instance_dict["port"], + } + }, + INSTALLED_APPS=[ + "django.contrib.postgres", + "django.contrib.contenttypes", + django_project, + ], + ) 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 + print("DJANGO_PROJECT", django_project) if "gr" in django_project: call_command("makemigrations", "common", interactive=False) + else: + call_command("makemigrations", interactive=False) call_command("migrate") diff --git a/test_utils/managers/gr/conftest.py b/test_utils/managers/gr/conftest.py index a7fa9e9..b5db2a5 100644 --- a/test_utils/managers/gr/conftest.py +++ b/test_utils/managers/gr/conftest.py @@ -24,6 +24,8 @@ if TYPE_CHECKING: # === Msc === + + @pytest.fixture(scope="session") def gr_redis_config_db() -> str: return str(randint(99, 1_023)) diff --git a/test_utils/managers/thl/conftest.py b/test_utils/managers/thl/conftest.py index 391b74c..6e19bef 100644 --- a/test_utils/managers/thl/conftest.py +++ b/test_utils/managers/thl/conftest.py @@ -45,20 +45,7 @@ if TYPE_CHECKING: WallManager, ) - -@pytest.fixture(scope="session") -def thl_web_rr(django_db_factory: Callable[..., PostgresDsn]) -> PostgresConfig: - - return PostgresConfig( - dsn=django_db_factory("generalresearch.thl_django"), - connect_timeout=1, - statement_timeout=5, - ) - - -@pytest.fixture(scope="session") -def thl_web_rw(thl_web_rr: PostgresConfig) -> PostgresConfig: - return thl_web_rr +# === Msc === @pytest.fixture(scope="session") @@ -97,6 +84,25 @@ def thl_redis_config( r.flushdb() +@pytest.fixture(scope="session") +def thl_web_rr(django_db_factory: Callable[..., PostgresDsn]) -> PostgresConfig: + _dsn = django_db_factory("generalresearch.thl_django") + + return PostgresConfig( + dsn=_dsn, + connect_timeout=1, + statement_timeout=5, + ) + + +@pytest.fixture(scope="session") +def thl_web_rw(thl_web_rr: PostgresConfig) -> PostgresConfig: + return thl_web_rr + + +# === Managers === + + @pytest.fixture(scope="session") def payout_event_manager( thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig diff --git a/test_utils/models/gr/conftest.py b/test_utils/models/gr/conftest.py index b87f3bb..a73dd70 100644 --- a/test_utils/models/gr/conftest.py +++ b/test_utils/models/gr/conftest.py @@ -66,35 +66,58 @@ def gr_user_cache( return gr_user +# --- Business Bank Account --- + + @pytest.fixture def gr_business_bank_account_factory( - gr_bbam: BusinessBankAccountManager, + gr_business_bank_account_manager: BusinessBankAccountManager, ) -> Callable[..., BusinessBankAccount]: def _inner( business_id: PositiveInt, + save: bool = True, uuid: UUIDStr | None = None, transfer_method: TransferMethod | None = None, account_number: str | None = None, routing_number: str | None = None, iban: str | None = None, swift: str | None = None, - ): - from generalresearch.models.gr.business import TransferMethod + **kwargs, + ) -> BusinessBankAccount: - return gr_bbam.create( - business_id=business_id, - uuid=uuid or uuid4().hex, - transfer_method=transfer_method or TransferMethod.ACH, - account_number=account_number or uuid4().hex[:6], - routing_number=routing_number or uuid4().hex[:6], - iban=iban or uuid4().hex[:6], - swift=swift or uuid4().hex[:6], - ) + if save: + return gr_business_bank_account_manager.create( + business_id=business_id, + uuid=uuid or uuid4().hex, + transfer_method=transfer_method or TransferMethod.ACH, + account_number=account_number or uuid4().hex[:6], + routing_number=routing_number or uuid4().hex[:6], + iban=iban or uuid4().hex[:6], + swift=swift or uuid4().hex[:6], + **kwargs, + ) + else: + raise ValueError("BusinessBankAccount Business not supported yet") return _inner +@pytest.fixture +def gr_business_bank_account(gr_business_factory: Callable[..., Business]) -> Business: + return gr_business_factory(save=True) + + +@pytest.fixture +def unsaved_gr_business_bank_account( + gr_business_factory: Callable[..., Business], +) -> Business: + return gr_business_factory(save=False) + + +# ----------------- + + @pytest.fixture def gr_business_address_factory( gr_bam: BusinessAddressManager, @@ -204,14 +227,6 @@ def business_address( return business_address_manager.create_dummy(business_id=gr_business.id) -@pytest.fixture -def business_bank_account( - gr_business: Business, - business_bank_account_manager: BusinessBankAccountManager, -) -> BusinessBankAccount: - return business_bank_account_manager.create_dummy(business_id=gr_business.id) - - @pytest.fixture() def gr_user_token_header(gr_user_token: GRToken) -> dict[str, str]: return gr_user_token.auth_header diff --git a/tests/managers/gr/test_business.py b/tests/managers/gr/test_business.py index 35c471e..3513af5 100644 --- a/tests/managers/gr/test_business.py +++ b/tests/managers/gr/test_business.py @@ -25,18 +25,18 @@ class TestBusinessBankAccountManager: def test_init( self, - business_bank_account_manager: BusinessBankAccountManager, + gr_business_bank_account_manager: BusinessBankAccountManager, gr_db: PostgresConfig, ): - assert business_bank_account_manager.pg_config == gr_db + assert gr_business_bank_account_manager.pg_config == gr_db def test_create( self, gr_business: Business, - business_bank_account_manager: BusinessBankAccountManager, + gr_business_bank_account_manager: BusinessBankAccountManager, ): - instance = business_bank_account_manager.create( + instance = gr_business_bank_account_manager.create( business_id=gr_business.id, uuid=uuid4().hex, transfer_method=TransferMethod.ACH, @@ -44,7 +44,7 @@ class TestBusinessBankAccountManager: assert isinstance(instance, BusinessBankAccount) assert isinstance(instance.id, int) - res = business_bank_account_manager.get_by_business_id( + res = gr_business_bank_account_manager.get_by_business_id( business_id=instance.business_id ) assert isinstance(res, list) diff --git a/tests/models/gr/test_business.py b/tests/models/gr/test_business.py index 90e69db..57f31f3 100644 --- a/tests/models/gr/test_business.py +++ b/tests/models/gr/test_business.py @@ -62,12 +62,12 @@ class TestBusinessBankAccount: def test_init( self, gr_business: Business, - business_bank_account_manager: BusinessBankAccountManager, + gr_business_bank_account_manager: BusinessBankAccountManager, ): from generalresearch.models.gr.business import BusinessBankAccount from generalresearch.models.gr.definitions import TransferMethod - instance = business_bank_account_manager.create( + instance = gr_business_bank_account_manager.create( business_id=gr_business.id, uuid=uuid4().hex, transfer_method=TransferMethod.ACH, @@ -76,20 +76,20 @@ class TestBusinessBankAccount: def test_business( self, - business_bank_account: BusinessBankAccount, + gr_business_bank_account: BusinessBankAccount, gr_business: Business, gr_db: PostgresConfig, gr_redis_config: RedisConfig, ): from generalresearch.models.gr.business import Business - assert business_bank_account.business is None + assert gr_business_bank_account.business is None - business_bank_account.prefetch_business( + gr_business_bank_account.prefetch_business( pg_config=gr_db, redis_config=gr_redis_config ) - assert isinstance(business_bank_account.business, Business) - assert business_bank_account.business.uuid == gr_business.uuid + assert isinstance(gr_business_bank_account.business, Business) + assert gr_business_bank_account.business.uuid == gr_business.uuid class TestBusinessAddress: @@ -264,13 +264,13 @@ class TestBusiness: def test_bank_accounts( self, gr_business: Business, - business_bank_account_manager: BusinessBankAccountManager, + gr_business_bank_account_manager: BusinessBankAccountManager, ): assert gr_business.products is None # It's an empty list after prefetch gr_business.prefetch_bank_accounts( - business_bank_account_manager=business_bank_account_manager + business_bank_account_manager=gr_business_bank_account_manager ) assert isinstance(gr_business.bank_accounts, list) assert len(gr_business.bank_accounts) == 1 @@ -423,7 +423,7 @@ class TestBusiness: def test_pop_financial( self, gr_business: Business, - thl_web_rr: PostgresConfig, + product_manager: ProductManager, thl_ledger_manager: ThlLedgerManager, mnt_filepath: GRLDatasets, client_no_amm: DaskClient, @@ -431,7 +431,7 @@ class TestBusiness: ): assert gr_business.pop_financial is None gr_business.prebuild_pop_financial( - thl_pg_config=thl_web_rr, + product_manager=product_manager, thl_lm=thl_ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -442,7 +442,6 @@ class TestBusiness: def test_bp_accounts( self, gr_business: Business, - thl_web_rr: PostgresConfig, product_factory: Callable[..., Product], thl_ledger_manager: ThlLedgerManager, product_manager: ProductManager, diff --git a/tests/test_postgres.py b/tests/test_postgres.py index c53f644..9794321 100644 --- a/tests/test_postgres.py +++ b/tests/test_postgres.py @@ -68,4 +68,15 @@ class TestPostgresDjangoCreation: WHERE table_schema = 'public'; """) assert len(res) == 1 - assert res[0]["count"] == 56 + assert res[0]["count"] == 57 + + def test_django_tables_with_gr( + self, thl_web_rw: PostgresConfig, gr_db: 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"] > 57 -- cgit v1.2.3 From 97b14e2f133bda76f548ec1a522d9582c657d736 Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Wed, 2 Sep 2026 16:51:06 -0700 Subject: managers/gr/test_auth is all green ✅ --- generalresearch/thl_django/app/test_settings.py | 2 +- test_utils/managers/upk/conftest.py | 8 +- test_utils/models/gr/conftest.py | 254 ++++++++++++++++-------- tests/managers/gr/test_authentication.py | 115 +++++++---- 4 files changed, 250 insertions(+), 129 deletions(-) (limited to 'test_utils/models') diff --git a/generalresearch/thl_django/app/test_settings.py b/generalresearch/thl_django/app/test_settings.py index c5df32a..2738aed 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-02-a0310b', + "NAME": 'unittest-2026-09-02-77ae16', "USER": 'jenkins', "PASSWORD": '123456789', "HOST": 'unittest-postgresql.fmt2.grl.internal', diff --git a/test_utils/managers/upk/conftest.py b/test_utils/managers/upk/conftest.py index f581278..23af1b3 100644 --- a/test_utils/managers/upk/conftest.py +++ b/test_utils/managers/upk/conftest.py @@ -13,11 +13,9 @@ from generalresearch.managers.thl.profiling.uqa import UQAManager from generalresearch.managers.thl.profiling.user_upk import ( UserUpkManager, ) - -if TYPE_CHECKING: - from generalresearch.models.thl.user import User - from generalresearch.pg_helper import PostgresConfig - from generalresearch.redis_helper import RedisConfig +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/gr/conftest.py b/test_utils/models/gr/conftest.py index a73dd70..a5abf74 100644 --- a/test_utils/models/gr/conftest.py +++ b/test_utils/models/gr/conftest.py @@ -36,36 +36,6 @@ if TYPE_CHECKING: # --- Factory / Database --- -@pytest.fixture -def gr_user_factory(gr_user_manager: GRUserManager) -> Callable[..., GRUser]: - - def _inner( - sub: str | None = None, - is_superuser: bool = False, - ) -> GRUser: - sub = sub or f"{uuid4().hex}-{uuid4().hex}" - - return gr_user_manager.create( - sub=sub, - is_superuser=is_superuser, - ) - - return _inner - - -@pytest.fixture -def gr_user_cache( - gr_user: GRUser, - gr_db: PostgresConfig, - thl_web_rr: PostgresConfig, - gr_redis_config: RedisConfig, -) -> GRUser: - gr_user.set_cache( - pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config - ) - return gr_user - - # --- Business Bank Account --- @@ -98,7 +68,7 @@ def gr_business_bank_account_factory( **kwargs, ) else: - raise ValueError("BusinessBankAccount Business not supported yet") + raise ValueError("Unsaved BusinessBankAccount not supported yet") return _inner @@ -115,16 +85,17 @@ def unsaved_gr_business_bank_account( return gr_business_factory(save=False) -# ----------------- +# --- Business Address --- @pytest.fixture def gr_business_address_factory( - gr_bam: BusinessAddressManager, + gr_business_address_manager: BusinessAddressManager, ) -> Callable[..., BusinessAddress]: def _inner( business_id: PositiveInt, + save: bool = True, uuid: UUIDStr | None = None, line_1: str | None = None, line_2: str | None = None, @@ -133,7 +104,7 @@ def gr_business_address_factory( postal_code: str | None = None, phone_number: PhoneNumber | None = None, country: str | None = None, - ): + ) -> BusinessAddress: uuid = uuid or uuid4().hex line_1 = line_1 or "abc" line_2 = line_2 or "bczx" @@ -143,21 +114,48 @@ def gr_business_address_factory( phone_number = None country = country or "US" - return gr_bam.create( - business_id=business_id, - uuid=uuid, - line_1=line_1, - line_2=line_2, - city=city, - state=state, - postal_code=postal_code, - phone_number=phone_number, - country=country, - ) + if save: + return gr_business_address_manager.create( + business_id=business_id, + uuid=uuid, + line_1=line_1, + line_2=line_2, + city=city, + state=state, + postal_code=postal_code, + phone_number=phone_number, + country=country, + ) + else: + raise ValueError("Unsaved BusinessAddress not supported yet") 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], +) -> BusinessAddress: + return gr_business_address_factory(save=True) + + +@pytest.fixture +def unsaved_gr_business_address( + gr_business_address_factory: Callable[..., BusinessAddress], +) -> BusinessAddress: + return gr_business_address_factory(save=False) + + +# --- Business --- + + @pytest.fixture def gr_business_factory( gr_business_manager: BusinessManager, @@ -194,37 +192,127 @@ def unsaved_gr_business(gr_business_factory: Callable[..., Business]) -> Busines return gr_business_factory(save=False) +# --- GR Team --- + + @pytest.fixture -def gr_team( - gr_tm: TeamManager, +def gr_team_factory( + gr_team_manager: TeamManager, ) -> Callable[..., Team]: - def _inner(uuid: UUIDStr | None = None, name: str | None = None) -> Team: - uuid = uuid or uuid4().hex - name = name or f"name-{uuid4().hex[:12]}" + def _inner( + save: bool = True, + uuid: UUIDStr | None = None, + name: str | None = None, + **kwargs, + ) -> Team: + + if save: + return gr_team_manager.create(uuid=uuid, name=name, **kwargs) - return gr_tm.create(uuid=uuid, name=name) + else: + raise ValueError("BusinessBankAccount Business not supported yet") return _inner -@pytest.fixture() -def gr_user_token( - gr_user: GRUser, gr_tm: GRTokenManager, gr_db: PostgresConfig -) -> GRToken: - gr_tm.create(user_id=gr_user.id) - gr_user.prefetch_token(pg_config=gr_db) +@pytest.fixture +def gr_team(gr_team_factory: Callable[..., Team]) -> Team: + return gr_team_factory(save=True) + + +@pytest.fixture +def unsaved_gr_team( + gr_team_factory: Callable[..., Team], +) -> Team: + return gr_team_factory(save=False) - res = gr_user.token - assert res is not None, "GRToken should exist after creation and prefetching" - return res + +# --- GR User --- @pytest.fixture -def business_address( - gr_business: Business, business_address_manager: BusinessAddressManager -) -> BusinessAddress: - return business_address_manager.create_dummy(business_id=gr_business.id) +def gr_user_factory(gr_user_manager: GRUserManager) -> Callable[..., GRUser]: + + def _inner( + save: bool = True, + sub: str | None = None, + is_superuser: bool = False, + ) -> GRUser: + sub = sub or f"{uuid4().hex}-{uuid4().hex}" + + if save: + return gr_user_manager.create( + sub=sub, + is_superuser=is_superuser, + ) + else: + raise ValueError("Unsaved GR User not supported yet") + + return _inner + + +@pytest.fixture +def gr_user_cache( + gr_user: GRUser, + gr_db: PostgresConfig, + thl_web_rr: PostgresConfig, + gr_redis_config: RedisConfig, +) -> GRUser: + gr_user.set_cache( + pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config + ) + return gr_user + + +@pytest.fixture +def gr_user(gr_user_factory: Callable[..., GRUser]) -> GRUser: + return gr_user_factory(save=True) + + +@pytest.fixture +def unsaved_gr_user( + gr_user_factory: Callable[..., GRUser], +) -> GRUser: + return gr_user_factory(save=False) + + +# --- GR User Token --- + + +@pytest.fixture +def gr_user_token_factory( + gr_user: GRUser, gr_user_token_manager: GRUser, gr_db: PostgresConfig +) -> Callable[..., GRToken]: + + def _inner( + save: bool = True, + ) -> GRToken: + + if save: + gr_user_token_manager.create(user_id=gr_user.id) + gr_user.prefetch_token(pg_config=gr_db) + + res = gr_user.token + assert ( + res is not None + ), "GRToken should exist after creation and prefetching" + return res + + else: + raise ValueError("Unsaved GR User not supported yet") + + return _inner + + +@pytest.fixture +def gr_user_token(gr_user_token_factory: Callable[..., GRToken]) -> GRToken: + return gr_user_token_factory(save=True) + + +@pytest.fixture +def unsaved_gr_user_token(gr_user_token_factory: Callable[..., GRToken]) -> GRToken: + return gr_user_token_factory(save=False) @pytest.fixture() @@ -232,26 +320,32 @@ def gr_user_token_header(gr_user_token: GRToken) -> dict[str, str]: return gr_user_token.auth_header -@pytest.fixture(scope="function") -def membership(team: Team, gr_user: GRUser, team_manager: TeamManager) -> Membership: - assert team.id, "Team must be saved" - assert gr_user.id, "GRUser must be saved" - return team_manager.add_user(team=team, gr_user=gr_user) +# --- GR Membership --- -@pytest.fixture(scope="function") -def membership_factory( - team: Team, +@pytest.fixture() +def gr_membership_factory( + gr_team: Team, gr_user: GRUser, - membership_manager: MembershipManager, - team_manager: TeamManager, - gr_um: GRUserManager, + gr_membership_manager: MembershipManager, ) -> Callable[..., Membership]: - def _inner(**kwargs) -> Membership: - _team = kwargs.get("team", team_manager.create_dummy()) - _gr_user = kwargs.get("gr_user", gr_um.create_dummy()) - - return membership_manager.create(team=_team, gr_user=_gr_user) + def _inner(save: bool = True, **kwargs) -> Membership: + if save: + return gr_membership_manager.create(team=gr_team, gr_user=gr_user, **kwargs) + else: + raise ValueError("Unsaved GR Membership not supported yet") return _inner + + +@pytest.fixture() +def gr_membership(gr_membership_factory: Callable[..., Membership]) -> Membership: + return gr_membership_factory(save=True) + + +@pytest.fixture() +def unsaved_gr_membership( + gr_membership_factory: Callable[..., Membership], +) -> Membership: + return gr_membership_factory(save=False) diff --git a/tests/managers/gr/test_authentication.py b/tests/managers/gr/test_authentication.py index b9f43a6..0bcabc5 100644 --- a/tests/managers/gr/test_authentication.py +++ b/tests/managers/gr/test_authentication.py @@ -1,117 +1,146 @@ import logging +from collections.abc import Callable from uuid import uuid4 import pytest -from generalresearch.models.gr.authentication import GRUser +from generalresearch.managers.gr.authentication import GRTokenManager, GRUserManager +from generalresearch.managers.gr.team import TeamManager +from generalresearch.models.gr.authentication import GRToken, GRUser +from generalresearch.pg_helper import PostgresConfig +from generalresearch.redis_helper import RedisConfig SSO_ISSUER = "" class TestGRUserManager: - def test_create(self, gr_um): - - user: GRUser = gr_um.create_dummy() - instance = gr_um.get_by_id(user.id) - assert user.id == instance.id + def test_create(self, gr_user: GRUser, gr_user_manager: GRUserManager): + instance = gr_user_manager.get_by_id(gr_user.id) + assert isinstance(instance, GRUser) + assert gr_user.id == instance.id - instance2 = gr_um.get_by_id(user.id) - assert user.model_dump_json() == instance2.model_dump_json() + instance2 = gr_user_manager.get_by_id(gr_user.id) + assert isinstance(instance2, GRUser) + assert gr_user.model_dump_json() == instance2.model_dump_json() - def test_get_by_id(self, gr_user, gr_um): + def test_get_by_id(self, gr_user: GRUser, gr_user_manager: GRUserManager): with pytest.raises(expected_exception=ValueError) as cm: - gr_um.get_by_id(gr_user_id=999_999_999) + gr_user_manager.get_by_id(gr_user_id=999_999_999) assert "GRUser not found" in str(cm.value) - instance = gr_um.get_by_id(gr_user_id=gr_user.id) + instance = gr_user_manager.get_by_id(gr_user_id=gr_user.id) + assert isinstance(instance, GRUser) assert instance.sub == gr_user.sub - def test_get_by_sub(self, gr_user, gr_um): + def test_get_by_sub(self, gr_user: GRUser, gr_user_manager: GRUserManager): with pytest.raises(expected_exception=ValueError) as cm: - gr_um.get_by_sub(sub=uuid4().hex) + gr_user_manager.get_by_sub(sub=uuid4().hex) assert "GRUser not found" in str(cm.value) - instance = gr_um.get_by_sub(sub=gr_user.sub) + instance = gr_user_manager.get_by_sub(sub=gr_user.sub) + assert isinstance(instance, GRUser) assert instance.id == gr_user.id - def test_get_by_sub_or_create(self, gr_user, gr_um): + def test_get_by_sub_or_create( + self, gr_user: GRUser, gr_user_manager: GRUserManager + ): sub = f"{uuid4().hex}-{uuid4().hex}" with pytest.raises(expected_exception=ValueError) as cm: - gr_um.get_by_sub(sub=sub) + gr_user_manager.get_by_sub(sub=sub) assert "GRUser not found" in str(cm.value) - instance = gr_um.get_by_sub_or_create(sub=sub) + instance = gr_user_manager.get_by_sub_or_create(sub=sub) assert isinstance(instance, GRUser) assert instance.sub == sub - def test_get_all(self, gr_um): - res1 = gr_um.get_all() + def test_get_all( + self, gr_user_factory: Callable[..., GRUser], gr_user_manager: GRUserManager + ): + res1 = gr_user_manager.get_all() assert isinstance(res1, list) - gr_um.create_dummy() - res2 = gr_um.get_all() + gr_user_factory(save=True) + res2 = gr_user_manager.get_all() assert len(res1) == len(res2) - 1 - def test_get_by_team(self, gr_um): - res = gr_um.get_by_team(team_id=999_999_999) + def test_get_by_team(self, gr_user_manager: GRUserManager): + res = gr_user_manager.get_by_team(team_id=999_999_999) assert isinstance(res, list) assert res == [] - def test_list_product_uuids(self, caplog, gr_user, gr_um, thl_web_rr): + def test_list_product_uuids( + self, + caplog, + gr_user: GRUser, + gr_user_manager: GRUserManager, + thl_web_rr: PostgresConfig, + ): with caplog.at_level(logging.WARNING): - gr_um.list_product_uuids(user=gr_user, thl_pg_config=thl_web_rr) + gr_user_manager.list_product_uuids(user=gr_user, thl_pg_config=thl_web_rr) assert "prefetch not run" in caplog.text class TestGRTokenManager: - def test_create(self, gr_user, gr_tm): - assert gr_tm.create(user_id=gr_user.id) is None + def test_create(self, gr_user: GRUser, gr_team_manager: TeamManager): + assert gr_team_manager.create(user_id=gr_user.id) is None - token = gr_tm.get_by_user_id(user_id=gr_user.id) + token = gr_team_manager.get_by_user_id(user_id=gr_user.id) assert gr_user.id == token.user_id - def test_get_by_user_id(self, gr_user, gr_tm): - assert gr_tm.create(user_id=gr_user.id) is None + def test_get_by_user_id(self, gr_user: GRUser, gr_team_manager: TeamManager): + assert gr_team_manager.create(user_id=gr_user.id) is None - token = gr_tm.get_by_user_id(user_id=gr_user.id) + token = gr_team_manager.get_by_user_id(user_id=gr_user.id) assert gr_user.id == token.user_id - def test_prefetch_user(self, gr_user, gr_tm, gr_db, gr_redis_config): - from generalresearch.models.gr.authentication import GRToken + def test_prefetch_user( + self, + gr_user: GRUser, + gr_team_manager: TeamManager, + gr_db: PostgresConfig, + gr_redis_config: RedisConfig, + ): - gr_tm.create(user_id=gr_user.id) + gr_team_manager.create(user_id=gr_user.id) - token: GRToken = gr_tm.get_by_user_id(user_id=gr_user.id) + token: GRToken = gr_team_manager.get_by_user_id(user_id=gr_user.id) assert token.user is None token.prefetch_user(pg_config=gr_db, redis_config=gr_redis_config) assert token.user.id == gr_user.id - def test_get_by_key(self, gr_user, gr_um, gr_tm): - gr_tm.create(user_id=gr_user.id) - token = gr_tm.get_by_user_id(user_id=gr_user.id) + def test_get_by_key( + self, + gr_user: GRUser, + gr_team_manager: TeamManager, + ): + gr_team_manager.create(user_id=gr_user.id) + token = gr_team_manager.get_by_user_id(user_id=gr_user.id) - instance = gr_tm.get_by_key(api_key=token.key) + instance = gr_team_manager.get_by_key(api_key=token.key) assert token.created == instance.created # Search for non-existent key with pytest.raises(expected_exception=Exception) as cm: - gr_tm.get_by_key(api_key=uuid4().hex) + gr_team_manager.get_by_key(api_key=uuid4().hex) assert "No GRUser with token of " in str(cm.value) @pytest.mark.skip(reason="no idea how to actually test this...") - def test_get_by_sso_key(self, gr_user, gr_um, gr_tm, gr_redis_config): - from generalresearch.models.gr.authentication import GRToken + def test_get_by_sso_key( + self, + gr_team_manager: TeamManager, + gr_redis_config: RedisConfig, + ): api_key = "..." jwks = { # ... } - instance = gr_tm.get_by_key( + instance = gr_team_manager.get_by_key( api_key=api_key, jwks=jwks, audience="...", -- cgit v1.2.3 From ad620d7586640534a092672b8f3cddf6eff5604b Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Wed, 2 Sep 2026 17:40:18 -0700 Subject: gr mangers all green ✅ --- Jenkinsfile | 16 +++++ generalresearch/managers/gr/team.py | 7 +- generalresearch/models/thl/product.py | 4 +- generalresearch/thl_django/app/test_settings.py | 2 +- test_utils/managers/gr/conftest.py | 22 +++++- test_utils/models/conftest.py | 31 +-------- test_utils/models/gr/conftest.py | 10 +-- test_utils/models/thl/conftest.py | 61 +++++++++++------ tests/managers/gr/test_authentication.py | 32 +++++---- tests/managers/gr/test_business.py | 22 ++++-- tests/managers/gr/test_team.py | 91 +++++++++++++++---------- 11 files changed, 180 insertions(+), 118 deletions(-) (limited to 'test_utils/models') diff --git a/Jenkinsfile b/Jenkinsfile index de909b2..a646d22 100644 --- a/Jenkinsfile +++ b/Jenkinsfile @@ -60,6 +60,14 @@ pipeline { } stage('base') { + steps { + dir("generalresearch-${VER}") { + sh "${VENV}-${VER}/bin/pytest tests/test_postgres.py -vs" + } + } + } + + stage('models') { steps { dir("generalresearch-${VER}") { sh "${VENV}-${VER}/bin/pytest tests/models/gr/test_base.py -vs" @@ -67,6 +75,14 @@ pipeline { } } + stage('managers') { + steps { + dir("generalresearch-${VER}") { + sh "${VENV}-${VER}/bin/pytest tests/managers/gr/ -vs" + } + } + } + } } } diff --git a/generalresearch/managers/gr/team.py b/generalresearch/managers/gr/team.py index e551f85..41af709 100644 --- a/generalresearch/managers/gr/team.py +++ b/generalresearch/managers/gr/team.py @@ -11,6 +11,7 @@ from generalresearch.managers.base import ( PostgresManager, PostgresManagerWithRedis, ) +from generalresearch.managers.gr.authentication import GRUserManager from generalresearch.models.custom_types import UUIDStr from generalresearch.models.gr.team import ( Membership, @@ -187,10 +188,12 @@ class TeamManager(PostgresManagerWithRedis): return team - def add_user(self, team: Team, gr_user: GRUser) -> Membership: + def add_user( + self, team: Team, gr_user: GRUser, gr_user_manager: GRUserManager + ) -> Membership: """Create a Membership between a GRUser and a Team""" - team.prefetch_gr_users(pg_config=self.pg_config, redis_config=self.redis_config) + team.prefetch_gr_users(gr_user_manager=gr_user_manager) assert gr_user not in team.gr_users, ( "Can't create multiple Memberships for " "the same User to the same Team" diff --git a/generalresearch/models/thl/product.py b/generalresearch/models/thl/product.py index 346a98b..3677ff2 100644 --- a/generalresearch/models/thl/product.py +++ b/generalresearch/models/thl/product.py @@ -1396,8 +1396,8 @@ class Product(BaseModel, validate_assignment=True): # --- ORM --- - def model_dump_mysql(self) -> dict[str, Any]: - d = self.model_dump(mode="json") + def model_dump_mysql(self, *args, **kwargs) -> dict[str, Any]: + d = self.model_dump(mode="json", *args, **kwargs) assert self.created if "created" in d: diff --git a/generalresearch/thl_django/app/test_settings.py b/generalresearch/thl_django/app/test_settings.py index 2738aed..276b94a 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-02-77ae16', + "NAME": 'unittest-2026-09-03-ab1271', "USER": 'jenkins', "PASSWORD": '123456789', "HOST": 'unittest-postgresql.fmt2.grl.internal', diff --git a/test_utils/managers/gr/conftest.py b/test_utils/managers/gr/conftest.py index b5db2a5..cc1053c 100644 --- a/test_utils/managers/gr/conftest.py +++ b/test_utils/managers/gr/conftest.py @@ -7,7 +7,6 @@ from typing import TYPE_CHECKING import pytest import redis -import redis.asyncio as redis_async from pydantic import PostgresDsn from generalresearch.managers.gr.business import ( @@ -15,12 +14,14 @@ from generalresearch.managers.gr.business import ( BusinessBankAccountManager, BusinessManager, ) +from generalresearch.managers.gr.team import MembershipManager 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 + from generalresearch.managers.gr.team import TeamManager # === Msc === @@ -89,7 +90,17 @@ def gr_user_manager( @pytest.fixture(scope="session") -def gr_team_manager(gr_db: PostgresConfig) -> GRTokenManager: +def gr_team_manager(gr_db: PostgresConfig, gr_redis_config: RedisConfig) -> TeamManager: + assert gr_db.dsn.path + assert "/unittest-" in gr_db.dsn.path + + from generalresearch.managers.gr.team import TeamManager + + return TeamManager(pg_config=gr_db, redis_config=gr_redis_config) + + +@pytest.fixture(scope="session") +def gr_token_manager(gr_db: PostgresConfig) -> GRTokenManager: assert gr_db.dsn.path assert "/unittest-" in gr_db.dsn.path @@ -117,3 +128,10 @@ def gr_business_address_manager( gr_db: PostgresConfig, ) -> BusinessAddressManager: return BusinessAddressManager(pg_config=gr_db) + + +@pytest.fixture(scope="session") +def gr_membership_manager( + gr_db: PostgresConfig, +) -> MembershipManager: + return MembershipManager(pg_config=gr_db) diff --git a/test_utils/models/conftest.py b/test_utils/models/conftest.py index ed4da08..d71593f 100644 --- a/test_utils/models/conftest.py +++ b/test_utils/models/conftest.py @@ -286,36 +286,7 @@ def session( return session -@pytest.fixture -def product(request: Request, product_manager: ProductManager) -> Product: - - team = getattr(request, "team", None) - business = getattr(request, "business", None) - - return product_manager.create_dummy( - team_id=team.uuid if team else None, - business_id=business.uuid if business else None, - ) - - -@pytest.fixture -def product_factory(product_manager: ProductManager) -> Callable[..., Product]: - - def _inner( - team: Team | None = None, - business: Business | None = None, - commission_pct: Decimal = Decimal("0.05"), - ) -> Product: - return product_manager.create_dummy( - team_id=team.uuid if team else None, - business_id=business.uuid if business else None, - commission_pct=commission_pct, - ) - - return _inner - - -@pytest.fixture +@pytest.fixture() def payout_config(request: Request) -> PayoutConfig: from generalresearch.models.thl.product import ( PayoutConfig, diff --git a/test_utils/models/gr/conftest.py b/test_utils/models/gr/conftest.py index a5abf74..3dd73a1 100644 --- a/test_utils/models/gr/conftest.py +++ b/test_utils/models/gr/conftest.py @@ -207,8 +207,10 @@ def gr_team_factory( **kwargs, ) -> Team: + name = name or f"" + if save: - return gr_team_manager.create(uuid=uuid, name=name, **kwargs) + return gr_team_manager.create(name=name, uuid=uuid, **kwargs) else: raise ValueError("BusinessBankAccount Business not supported yet") @@ -325,12 +327,12 @@ def gr_user_token_header(gr_user_token: GRToken) -> dict[str, str]: @pytest.fixture() def gr_membership_factory( - gr_team: Team, - gr_user: GRUser, gr_membership_manager: MembershipManager, ) -> Callable[..., Membership]: - def _inner(save: bool = True, **kwargs) -> Membership: + def _inner( + gr_team: Team, gr_user: GRUser, save: bool = True, **kwargs + ) -> Membership: if save: return gr_membership_manager.create(team=gr_team, gr_user=gr_user, **kwargs) else: diff --git a/test_utils/models/thl/conftest.py b/test_utils/models/thl/conftest.py index 3545509..badd87c 100644 --- a/test_utils/models/thl/conftest.py +++ b/test_utils/models/thl/conftest.py @@ -39,6 +39,7 @@ if TYPE_CHECKING: from generalresearch.managers.thl.userhealth import AuditLogManager, IPRecordManager from generalresearch.managers.thl.wall import WallManager from generalresearch.models.definitions import DeviceType + 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 @@ -144,12 +145,18 @@ def wall_factory( return _inner -@pytest.fixture +# --- Product --- + + +@pytest.fixture() def product_factory(product_manager: ProductManager) -> Callable[..., Product]: def _inner( + save: bool = True, + team: Team | None = None, + # business: Business | None = None, + # commission_pct: Decimal = Decimal("0.05"), product_id: UUIDStr | None = None, - team_id: UUIDStr | None = None, business_id: UUIDStr | None = None, name: str | None = None, redirect_url: str | None = None, @@ -165,30 +172,46 @@ def product_factory(product_manager: ProductManager) -> Callable[..., Product]: ) -> Product: """To be used in tests, where we don't care about certain fields""" product_id = product_id if product_id else uuid4().hex - team_id = team_id if team_id else uuid4().hex + team_id = team.uuid if team else uuid4().hex name = name if name else f"name-{product_id[:12]}" redirect_url = redirect_url if redirect_url else "https://www.example.com/" - return product_manager.create( - product_id=product_id, - team_id=team_id, - business_id=business_id, - name=name, - redirect_url=redirect_url, - harmonizer_domain=harmonizer_domain, - commission_pct=commission_pct, - sources_config=sources_config, - payout_config=payout_config, - session_config=session_config, - profiling_config=profiling_config, - user_wallet_config=user_wallet_config, - user_create_config=user_create_config, - user_health_config=user_health_config, - ) + if save: + return product_manager.create( + product_id=product_id, + team_id=team_id, + business_id=business_id, + name=name, + redirect_url=redirect_url, + harmonizer_domain=harmonizer_domain, + commission_pct=commission_pct, + sources_config=sources_config, + payout_config=payout_config, + session_config=session_config, + profiling_config=profiling_config, + user_wallet_config=user_wallet_config, + user_create_config=user_create_config, + user_health_config=user_health_config, + ) + else: + raise ValueError("Unsaved Product not yet supported") return _inner +@pytest.fixture() +def product(product_factory: Callable[..., Product]) -> Product: + return product_factory(save=True) + + +@pytest.fixture() +def unsaved_product(product_factory: Callable[..., Product]) -> Product: + return product_factory(save=False) + + +# --- Session --- + + @pytest.fixture def session_factory(session_manager: SessionManager): diff --git a/tests/managers/gr/test_authentication.py b/tests/managers/gr/test_authentication.py index 0bcabc5..1310c79 100644 --- a/tests/managers/gr/test_authentication.py +++ b/tests/managers/gr/test_authentication.py @@ -84,29 +84,32 @@ class TestGRUserManager: class TestGRTokenManager: - def test_create(self, gr_user: GRUser, gr_team_manager: TeamManager): - assert gr_team_manager.create(user_id=gr_user.id) is None + def test_create(self, gr_user: GRUser, gr_token_manager: GRTokenManager): + assert gr_token_manager.create(user_id=gr_user.id) is None - token = gr_team_manager.get_by_user_id(user_id=gr_user.id) + token = gr_token_manager.get_by_user_id(user_id=gr_user.id) + assert isinstance(token, GRToken) assert gr_user.id == token.user_id - def test_get_by_user_id(self, gr_user: GRUser, gr_team_manager: TeamManager): - assert gr_team_manager.create(user_id=gr_user.id) is None + def test_get_by_user_id(self, gr_user: GRUser, gr_token_manager: GRTokenManager): + assert gr_token_manager.create(user_id=gr_user.id) is None - token = gr_team_manager.get_by_user_id(user_id=gr_user.id) + token = gr_token_manager.get_by_user_id(user_id=gr_user.id) + assert isinstance(token, GRToken) assert gr_user.id == token.user_id def test_prefetch_user( self, gr_user: GRUser, - gr_team_manager: TeamManager, + gr_token_manager: GRTokenManager, gr_db: PostgresConfig, gr_redis_config: RedisConfig, ): - gr_team_manager.create(user_id=gr_user.id) + gr_token_manager.create(user_id=gr_user.id) - token: GRToken = gr_team_manager.get_by_user_id(user_id=gr_user.id) + token: GRToken | None = gr_token_manager.get_by_user_id(user_id=gr_user.id) + assert isinstance(token, GRToken) assert token.user is None token.prefetch_user(pg_config=gr_db, redis_config=gr_redis_config) @@ -115,17 +118,18 @@ class TestGRTokenManager: def test_get_by_key( self, gr_user: GRUser, - gr_team_manager: TeamManager, + gr_token_manager: GRTokenManager, ): - gr_team_manager.create(user_id=gr_user.id) - token = gr_team_manager.get_by_user_id(user_id=gr_user.id) + gr_token_manager.create(user_id=gr_user.id) + token = gr_token_manager.get_by_user_id(user_id=gr_user.id) + assert isinstance(token, GRToken) - instance = gr_team_manager.get_by_key(api_key=token.key) + instance = gr_token_manager.get_by_key(api_key=token.key) assert token.created == instance.created # Search for non-existent key with pytest.raises(expected_exception=Exception) as cm: - gr_team_manager.get_by_key(api_key=uuid4().hex) + gr_token_manager.get_by_key(api_key=uuid4().hex) assert "No GRUser with token of " in str(cm.value) @pytest.mark.skip(reason="no idea how to actually test this...") diff --git a/tests/managers/gr/test_business.py b/tests/managers/gr/test_business.py index 3513af5..0d5b0d5 100644 --- a/tests/managers/gr/test_business.py +++ b/tests/managers/gr/test_business.py @@ -1,3 +1,4 @@ +from collections.abc import Callable from typing import TYPE_CHECKING from uuid import uuid4 @@ -9,6 +10,7 @@ from generalresearch.models.gr.business import ( BusinessBankAccount, ) from generalresearch.models.gr.definitions import TransferMethod +from generalresearch.models.gr.team import Team if TYPE_CHECKING: from generalresearch.managers.gr.business import ( @@ -68,9 +70,9 @@ class TestBusinessAddressManager: class TestBusinessManager: - def test_create(self, business_manager: BusinessManager): + def test_create(self, gr_business_factory: Callable[..., Business]): - instance = business_manager.create_dummy() + instance = gr_business_factory() assert isinstance(instance, Business) assert isinstance(instance.id, int) @@ -88,11 +90,15 @@ class TestBusinessManager: assert isinstance(res, Business) assert res.id == instance.id - def test_get_all(self, business_manager: BusinessManager): + def test_get_all( + self, + business_manager: BusinessManager, + gr_business_factory: Callable[..., Business], + ): res1 = business_manager.get_all() assert isinstance(res1, list) - business_manager.create_dummy() + gr_business_factory() res2 = business_manager.get_all() assert len(res1) == len(res2) - 1 @@ -106,17 +112,19 @@ class TestBusinessManager: gr_user: GRUser, team_manager: TeamManager, membership_manager: MembershipManager, + gr_business_factory: Callable[..., Business], + gr_team_factory: Callable[..., Team], ): res = business_manager.get_by_user_id(user_id=gr_user.id) assert len(res) == 0 # Create a business: Business, but don't add it to anything - b1 = business_manager.create_dummy() + b1 = gr_business_factory() res = business_manager.get_by_user_id(user_id=gr_user.id) assert len(res) == 0 # Create a Team, but don't create any Memberships - t1 = team_manager.create_dummy() + t1 = gr_team_factory() res = business_manager.get_by_user_id(user_id=gr_user.id) assert len(res) == 0 @@ -133,7 +141,7 @@ class TestBusinessManager: assert len(res) == 1 # Add another Business to the Team! - b2 = business_manager.create_dummy() + b2 = gr_business_factory() team_manager.add_business(team=t1, business=b2) res = business_manager.get_by_user_id(user_id=gr_user.id) assert len(res) == 2 diff --git a/tests/managers/gr/test_team.py b/tests/managers/gr/test_team.py index 17e0470..751e33c 100644 --- a/tests/managers/gr/test_team.py +++ b/tests/managers/gr/test_team.py @@ -4,6 +4,7 @@ from collections.abc import Callable from typing import TYPE_CHECKING from uuid import uuid4 +from generalresearch.models.gr.authentication import GRUser from generalresearch.models.gr.team import Membership, Team if TYPE_CHECKING: @@ -23,94 +24,110 @@ class TestMembershipManager: class TestTeamManager: - def test_init(self, team_manager: TeamManager, gr_db: PostgresConfig): - assert team_manager.pg_config == gr_db + def test_init(self, gr_team_manager: TeamManager, gr_db: PostgresConfig): + assert gr_team_manager.pg_config == gr_db - def test_get_or_create(self, team_manager: TeamManager): + def test_get_or_create(self, gr_team_manager: TeamManager): from generalresearch.models.gr.team import Team new_uuid = uuid4().hex - team: Team = team_manager.get_or_create(uuid=new_uuid) + team: Team = gr_team_manager.get_or_create(uuid=new_uuid) assert isinstance(team, Team) assert isinstance(team.id, int) assert team.uuid == new_uuid assert team.name == "< Unknown >" - def test_get_all(self, team_manager: TeamManager): - res1 = team_manager.get_all() + def test_get_all( + self, gr_team_factory: Callable[..., Team], gr_team_manager: TeamManager + ): + res1 = gr_team_manager.get_all() assert isinstance(res1, list) - team_manager.create_dummy() - res2 = team_manager.get_all() + gr_team_factory() + res2 = gr_team_manager.get_all() assert len(res1) == len(res2) - 1 - def test_create(self, team_manager: TeamManager): + def test_create( + self, gr_team_factory: Callable[..., Team], gr_team_manager: TeamManager + ): - team: Team = team_manager.create_dummy() + team: Team = gr_team_factory() assert isinstance(team, Team) assert isinstance(team.id, int) def test_add_user( self, - team: Team, - team_manager: TeamManager, - gr_um: GRUserManager, - gr_db: PostgresConfig, - gr_redis_config: RedisConfig, + gr_team: Team, + gr_team_manager: TeamManager, + gr_user_manager: GRUserManager, + gr_user_factory: Callable[..., GRUser], ): - user: GRUser = gr_um.create_dummy() + user: GRUser = gr_user_factory() - instance = team_manager.add_user(team=team, gr_user=user) + instance = gr_team_manager.add_user( + gr_user_manager=gr_user_manager, team=gr_team, gr_user=user + ) assert isinstance(instance, Membership) # assert team.gr_users is None - team.prefetch_gr_users(pg_config=gr_db, redis_config=gr_redis_config) - assert isinstance(team.gr_users, list) - assert len(team.gr_users) - assert team.gr_users == [user] + gr_team.prefetch_gr_users(gr_user_manager=gr_user_manager) + assert isinstance(gr_team.gr_users, list) + assert len(gr_team.gr_users) + assert gr_team.gr_users == [user] - def test_get_by_uuid(self, team_manager: TeamManager): + def test_get_by_uuid( + self, gr_team_factory: Callable[..., Team], gr_team_manager: TeamManager + ): - team: Team = team_manager.create_dummy() + team: Team = gr_team_factory() - instance = team_manager.get_by_uuid(team_uuid=team.uuid) + instance = gr_team_manager.get_by_uuid(team_uuid=team.uuid) + assert isinstance(instance, Team) assert team.id == instance.id - def test_get_by_id(self, team_manager: TeamManager): + def test_get_by_id( + self, gr_team_factory: Callable[..., Team], gr_team_manager: TeamManager + ): - team: Team = team_manager.create_dummy() + team: Team = gr_team_factory() - instance = team_manager.get_by_id(team_id=team.id) + instance = gr_team_manager.get_by_id(team_id=team.id) + assert isinstance(instance, Team) assert team.uuid == instance.uuid def test_get_by_user( - self, team: Team, team_manager: TeamManager, gr_um: GRUserManager + self, + gr_team: Team, + gr_user_factory: Callable[..., GRUser], + gr_team_manager: TeamManager, + gr_user_manager: GRUserManager, ): + user: GRUser = gr_user_factory() + gr_team_manager.add_user( + gr_user_manager=gr_user_manager, team=gr_team, gr_user=user + ) - user: GRUser = gr_um.create_dummy() - team_manager.add_user(team=team, gr_user=user) - - res = team_manager.get_by_user(gr_user=user) + res = gr_team_manager.get_by_user(gr_user=user) assert isinstance(res, list) assert len(res) == 1 instance = res[0] assert isinstance(instance, Team) - assert instance.uuid == team.uuid + assert instance.uuid == gr_team.uuid def test_get_by_user_duplicates( self, gr_user: GRUser, product_factory: Callable[..., Product], - membership_factory: Callable[..., Membership], - team: Team, + gr_membership_factory: Callable[..., Membership], + gr_team: Team, gr_redis_config: RedisConfig, gr_db: PostgresConfig, ): - product_factory(team=team) - membership_factory(team=team, gr_user=gr_user) + product_factory(team=gr_team) + gr_membership_factory(gr_team=gr_team, gr_user=gr_user) gr_user.prefetch_teams( pg_config=gr_db, -- cgit v1.2.3 From 4a4e5293777aa0fe5fb5bee4e61b1c6f735b4c13 Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Wed, 2 Sep 2026 18:29:22 -0700 Subject: WIP tests/models/gr --- generalresearch/models/gr/business.py | 37 ++++-- generalresearch/models/gr/team.py | 25 ++-- generalresearch/thl_django/app/test_settings.py | 2 +- test_utils/managers/conftest.py | 75 +---------- test_utils/managers/thl/conftest.py | 13 +- test_utils/models/thl/conftest.py | 8 +- tests/models/gr/test_base.py | 7 +- tests/models/gr/test_business.py | 57 +++++--- tests/models/gr/test_team.py | 169 ++++++++++++------------ 9 files changed, 185 insertions(+), 208 deletions(-) (limited to 'test_utils/models') diff --git a/generalresearch/models/gr/business.py b/generalresearch/models/gr/business.py index 73a2f27..2104650 100644 --- a/generalresearch/models/gr/business.py +++ b/generalresearch/models/gr/business.py @@ -302,7 +302,10 @@ class Business(BaseModel): # of knowing if a new Product has been added since the last time it # ran. self.prefetch_products(product_manager=product_manager) + assert isinstance(self.products, list) product_lookup = {p.uuid: p for p in self.products} + assert isinstance(self.product_uuids, list) + assert thl_lm.currency accounts = thl_lm.get_accounts_if_exists( qualified_names=[ @@ -339,7 +342,7 @@ class Business(BaseModel): def prebuild_balance( self, - thl_pg_config: PostgresConfig, + product_manager: ProductManager, lm: LedgerManager, ds: GRLDatasets, client: DaskClient, @@ -369,8 +372,9 @@ class Business(BaseModel): volume levels. """ LOG.debug(f"Business.prebuild_balance({self.uuid=})") + assert lm.currency - self.prefetch_products(thl_pg_config=thl_pg_config) + self.prefetch_products(product_manager=product_manager) accounts: list[LedgerAccount] = lm.get_accounts_if_exists( qualified_names=( @@ -425,7 +429,7 @@ class Business(BaseModel): df = df.groupby("account_id").sum() self.balance = BusinessBalances.from_pandas( - input_data=df, accounts=accounts, thl_pg_config=thl_pg_config + input_data=df, accounts=accounts, product_manager=product_manager ) return @@ -462,6 +466,7 @@ class Business(BaseModel): """ if self.bp_accounts is None: self.prefetch_bp_accounts(thl_lm=thl_lm, product_manager=product_manager) + assert isinstance(self.bp_accounts, list) from generalresearch.models.admin.request import ( ReportRequest, @@ -504,13 +509,13 @@ class Business(BaseModel): def prebuild_enriched_session_parquet( self, - thl_pg_config: PostgresConfig, + product_manager: ProductManager, ds: GRLDatasets, client: DaskClient, mnt_gr_api: Path, enriched_session: EnrichedSessionMerge | None = None, ) -> None: - self.prefetch_products(thl_pg_config=thl_pg_config) + self.prefetch_products(product_manager=product_manager) if enriched_session is None: from generalresearch.incite.defaults import ( @@ -547,13 +552,13 @@ class Business(BaseModel): def prebuild_enriched_wall_parquet( self, - thl_pg_config: PostgresConfig, + product_manager: ProductManager, ds: GRLDatasets, client: DaskClient, mnt_gr_api: Path, enriched_wall: EnrichedWallMerge | None = None, ) -> None: - self.prefetch_products(thl_pg_config=thl_pg_config) + self.prefetch_products(product_manager=product_manager) if enriched_wall is None: from generalresearch.incite.defaults import ( @@ -619,6 +624,8 @@ class Business(BaseModel): def set_cache( self, pg_config: PostgresConfig, + product_manager: ProductManager, + business_bank_account_manager: BusinessBankAccountManager, thl_web_rr: PostgresConfig, redis_config: RedisConfig, client: DaskClient, @@ -637,12 +644,14 @@ class Business(BaseModel): self.prefetch_addresses(pg_config=pg_config) self.prefetch_teams(pg_config=pg_config) - self.prefetch_products(thl_pg_config=thl_web_rr) - self.prefetch_bank_accounts(pg_config=pg_config) - self.prefetch_bp_accounts(thl_lm=thl_lm, thl_pg_config=thl_web_rr) + self.prefetch_products(product_manager=product_manager) + self.prefetch_bank_accounts( + business_bank_account_manager=business_bank_account_manager + ) + self.prefetch_bp_accounts(thl_lm=thl_lm, product_manager=product_manager) self.prebuild_balance( - thl_pg_config=thl_web_rr, + product_manager=product_manager, lm=lm, ds=ds, client=client, @@ -650,7 +659,7 @@ class Business(BaseModel): ) self.prebuild_payouts(bpem=bpem) self.prebuild_pop_financial( - thl_pg_config=thl_web_rr, + product_manager=product_manager, thl_lm=thl_lm, ds=ds, client=client, @@ -682,7 +691,7 @@ class Business(BaseModel): enriched_session = es(ds=ds) self.prebuild_enriched_session_parquet( - thl_pg_config=thl_web_rr, + product_manager=product_manager, client=client, ds=ds, mnt_gr_api=mnt_gr_api, @@ -695,7 +704,7 @@ class Business(BaseModel): enriched_wall = ew(ds=ds) self.prebuild_enriched_wall_parquet( - thl_pg_config=thl_web_rr, + product_manager=product_manager, client=client, ds=ds, mnt_gr_api=mnt_gr_api, diff --git a/generalresearch/models/gr/team.py b/generalresearch/models/gr/team.py index b1553e2..aa62c5a 100644 --- a/generalresearch/models/gr/team.py +++ b/generalresearch/models/gr/team.py @@ -142,13 +142,13 @@ class Team(BaseModel): def prebuild_enriched_session_parquet( self, - thl_pg_config: PostgresConfig, + product_manager: ProductManager, ds: GRLDatasets, client: Client, mnt_gr_api: Path, enriched_session: EnrichedSessionMerge | None = None, ) -> None: - self.prefetch_products(thl_pg_config=thl_pg_config) + self.prefetch_products(product_manager=product_manager) if enriched_session is None: from generalresearch.incite.defaults import ( @@ -185,13 +185,13 @@ class Team(BaseModel): def prebuild_enriched_wall_parquet( self, - thl_pg_config: PostgresConfig, + product_manager: ProductManager, ds: GRLDatasets, client: Client, mnt_gr_api: Path, enriched_wall: EnrichedWallMerge | None = None, ) -> None: - self.prefetch_products(thl_pg_config=thl_pg_config) + self.prefetch_products(product_manager=product_manager) if enriched_wall is None: from generalresearch.incite.defaults import ( @@ -259,7 +259,10 @@ class Team(BaseModel): def set_cache( self, - pg_config: PostgresConfig, + product_manager: ProductManager, + gr_user_manager: GRUserManager, + gr_business_manager: BusinessManager, + gr_membership_manager: MembershipManager, thl_web_rr: PostgresConfig, redis_config: RedisConfig, client: Client, @@ -268,10 +271,10 @@ class Team(BaseModel): enriched_session: EnrichedSessionMerge | None = None, enriched_wall: EnrichedWallMerge | None = None, ) -> None: - self.prefetch_products(thl_pg_config=thl_web_rr) - self.prefetch_gr_users(pg_config=pg_config, redis_config=redis_config) - self.prefetch_businesses(pg_config=pg_config, redis_config=redis_config) - self.prefetch_memberships(pg_config=pg_config) + self.prefetch_products(product_manager=product_manager) + self.prefetch_gr_users(gr_user_manager=gr_user_manager) + self.prefetch_businesses(business_manager=gr_business_manager) + self.prefetch_memberships(membership_manager=gr_membership_manager) rc = redis_config.create_redis_client() mapping = self.model_dump(mode="json") @@ -288,7 +291,7 @@ class Team(BaseModel): enriched_session = es(ds=ds) self.prebuild_enriched_session_parquet( - thl_pg_config=thl_web_rr, + product_manager=product_manager, client=client, ds=ds, mnt_gr_api=mnt_gr_api, @@ -301,7 +304,7 @@ class Team(BaseModel): enriched_wall = ew(ds=ds) self.prebuild_enriched_wall_parquet( - thl_pg_config=thl_web_rr, + product_manager=product_manager, client=client, ds=ds, mnt_gr_api=mnt_gr_api, diff --git a/generalresearch/thl_django/app/test_settings.py b/generalresearch/thl_django/app/test_settings.py index c9a1955..f3d23af 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-7aa1e5', + "NAME": 'unittest-2026-09-03-728bcf', "USER": 'jenkins', "PASSWORD": '123456789', "HOST": 'unittest-postgresql.fmt2.grl.internal', diff --git a/test_utils/managers/conftest.py b/test_utils/managers/conftest.py index ed771c7..ff088c2 100644 --- a/test_utils/managers/conftest.py +++ b/test_utils/managers/conftest.py @@ -14,15 +14,6 @@ from generalresearch.managers.thl.user_streak import ( 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 ( @@ -35,6 +26,7 @@ if TYPE_CHECKING: IPRecordManager, UserIpHistoryManager, ) + from generalresearch.models.thl.user import User from generalresearch.models.thl.wallet.cashout_method import CashoutMethod from generalresearch.pg_helper import PostgresConfig from generalresearch.redis_helper import RedisConfig @@ -100,7 +92,7 @@ def user_iphistory_manager( @pytest.fixture(scope="function") -def user_iphistory_manager_clear_cache(user_iphistory_manager, user): +def user_iphistory_manager_clear_cache(user_iphistory_manager, user: User): # On successive py-test/jenkins runs, the cache may contain # the previous run's info (keyed under the same user_id) user_iphistory_manager.delete_user_ip_history_cache(user_id=user.user_id) @@ -205,69 +197,6 @@ def spectrum_survey_manager(spectrum_rw: SqlHelper) -> SpectrumSurveyManager: return SpectrumSurveyManager(sql_helper=spectrum_rw) -# === GR === -@pytest.fixture(scope="session") -def business_manager( - gr_db: PostgresConfig, gr_redis_config: RedisConfig -) -> BusinessManager: - from generalresearch.redis_helper import RedisConfig - - assert gr_db.dsn.path - assert "/unittest-" in gr_db.dsn.path - assert isinstance(gr_redis_config, RedisConfig) - - from generalresearch.managers.gr.business import BusinessManager - - return BusinessManager( - pg_config=gr_db, - redis_config=gr_redis_config, - ) - - -@pytest.fixture(scope="session") -def business_address_manager(gr_db: PostgresConfig) -> BusinessAddressManager: - assert gr_db.dsn.path - assert "/unittest-" in gr_db.dsn.path - - from generalresearch.managers.gr.business import BusinessAddressManager - - return BusinessAddressManager(pg_config=gr_db) - - -@pytest.fixture(scope="session") -def business_bank_account_manager( - gr_db: PostgresConfig, -) -> BusinessBankAccountManager: - assert gr_db.dsn.path - assert "/unittest-" in gr_db.dsn.path - - from generalresearch.managers.gr.business import ( - BusinessBankAccountManager, - ) - - return BusinessBankAccountManager(pg_config=gr_db) - - -@pytest.fixture(scope="session") -def team_manager(gr_db: PostgresConfig, gr_redis_config: RedisConfig) -> TeamManager: - assert gr_db.dsn.path - assert "/unittest-" in gr_db.dsn.path - - from generalresearch.managers.gr.team import TeamManager - - return TeamManager(pg_config=gr_db, redis_config=gr_redis_config) - - -@pytest.fixture(scope="session") -def membership_manager(gr_db: PostgresConfig) -> MembershipManager: - assert gr_db.dsn.path - assert "/unittest-" in gr_db.dsn.path - - from generalresearch.managers.gr.team import MembershipManager - - return MembershipManager(pg_config=gr_db) - - @pytest.fixture(scope="session") def delete_buyers_surveys( thl_web_rw: PostgresConfig, buyer_manager: BuyerManager diff --git a/test_utils/managers/thl/conftest.py b/test_utils/managers/thl/conftest.py index 6e19bef..18a31e2 100644 --- a/test_utils/managers/thl/conftest.py +++ b/test_utils/managers/thl/conftest.py @@ -184,7 +184,10 @@ def product_manager(thl_web_rw: PostgresConfig) -> ProductManager: @pytest.fixture(scope="session") def user_manager( - settings: GRLBaseSettings, thl_web_rw: PostgresConfig, thl_web_rr: PostgresConfig + settings: GRLBaseSettings, + thl_web_rw: PostgresConfig, + thl_web_rr: PostgresConfig, + thl_redis_config: RedisConfig, ) -> UserManager: assert thl_web_rw.dsn assert thl_web_rw.dsn.path @@ -193,16 +196,22 @@ def user_manager( assert "/unittest-" in thl_web_rw.dsn.path assert "/unittest-" in thl_web_rr.dsn.path + from generalresearch.managers.thl.user_manager.rate_limit import UserManagerLimiter from generalresearch.managers.thl.user_manager.user_manager import ( UserManager, ) - return UserManager( + um = UserManager( pg_config=thl_web_rw, pg_config_rr=thl_web_rr, redis=settings.redis, ) + # rc = thl_redis_config.create_redis_client() + um.user_manager_limiter = UserManagerLimiter(redis=thl_redis_config.dsn) + + return um + @pytest.fixture(scope="session") def mysql_user_manager(thl_web_rw: PostgresConfig) -> MysqlUserManager: diff --git a/test_utils/models/thl/conftest.py b/test_utils/models/thl/conftest.py index badd87c..5826f0d 100644 --- a/test_utils/models/thl/conftest.py +++ b/test_utils/models/thl/conftest.py @@ -39,6 +39,7 @@ if TYPE_CHECKING: from generalresearch.managers.thl.userhealth import AuditLogManager, IPRecordManager from generalresearch.managers.thl.wall import WallManager 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 @@ -154,8 +155,7 @@ def product_factory(product_manager: ProductManager) -> Callable[..., Product]: def _inner( save: bool = True, team: Team | None = None, - # business: Business | None = None, - # commission_pct: Decimal = Decimal("0.05"), + business: Business | None = None, product_id: UUIDStr | None = None, business_id: UUIDStr | None = None, name: str | None = None, @@ -171,8 +171,12 @@ def product_factory(product_manager: ProductManager) -> Callable[..., Product]: user_health_config: UserHealthConfig | None = None, ) -> Product: """To be used in tests, where we don't care about certain fields""" + product_id = product_id if product_id else uuid4().hex + team_id = team.uuid if team else uuid4().hex + business_id = business.uuid if business else uuid4().hex + name = name if name else f"name-{product_id[:12]}" redirect_url = redirect_url if redirect_url else "https://www.example.com/" diff --git a/tests/models/gr/test_base.py b/tests/models/gr/test_base.py index fba0960..56603bf 100644 --- a/tests/models/gr/test_base.py +++ b/tests/models/gr/test_base.py @@ -41,10 +41,15 @@ class TestGRPostgresDjangoCreation: assert isinstance(dsn, PostgresDsn) def test_django_tables(self, gr_db: PostgresConfig): + """ + WARNING: This will always be the thl_django tables in addition + to the GR tables due to the way our fixtures are loaded. + """ + res = gr_db.execute_sql_query(query=""" SELECT COUNT(*) FROM information_schema.tables WHERE table_schema = 'public'; """) assert len(res) == 1 - assert res[0]["count"] == 10 + assert res[0]["count"] == 65 diff --git a/tests/models/gr/test_business.py b/tests/models/gr/test_business.py index 57f31f3..4e0b4e1 100644 --- a/tests/models/gr/test_business.py +++ b/tests/models/gr/test_business.py @@ -194,7 +194,7 @@ class TestBusiness: bpem=business_payout_event_manager, ) gr_business.prebuild_balance( - thl_pg_config=thl_web_rr, + product_manager=product_manager, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -283,12 +283,13 @@ class TestBusiness: thl_web_rr: PostgresConfig, ledger_manager: LedgerManager, pop_ledger_merge: PopLedgerMerge, + product_manager: ProductManager, ): assert gr_business.balance is None with pytest.raises(expected_exception=AssertionError) as cm: gr_business.prebuild_balance( - thl_pg_config=thl_web_rr, + product_manager=product_manager, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -409,8 +410,6 @@ class TestBusiness: ) gr_business.prebuild_payouts( - thl_pg_config=thl_web_rr, - thl_lm=thl_ledger_manager, bpem=business_payout_event_manager, ) @@ -495,6 +494,7 @@ class TestBusinessBalance: create_main_accounts: Callable[..., None], client_no_amm: DaskClient, ledger_collection, + product_manager: ProductManager, pop_ledger_merge: PopLedgerMerge, delete_df_collection: Callable[..., None], ): @@ -522,7 +522,7 @@ class TestBusinessBalance: 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, @@ -549,6 +549,7 @@ class TestBusinessBalance: user_factory: Callable[..., User], mnt_filepath: GRLDatasets, ledger_manager: LedgerManager, + product_manager: ProductManager, start: datetime, thl_web_rr: PostgresConfig, session_with_tx_factory: Callable[..., Session], @@ -582,7 +583,7 @@ class TestBusinessBalance: 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, @@ -631,6 +632,7 @@ class TestBusinessBalance: gr_business: Business, product_factory: Callable[..., Product], user_factory: Callable[..., User], + product_manager: ProductManager, mnt_filepath: GRLDatasets, bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], thl_ledger_manager: ThlLedgerManager, @@ -687,7 +689,7 @@ class TestBusinessBalance: 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, @@ -714,6 +716,7 @@ class TestBusinessBalance: payout_event_manager: PayoutEventManager, session_with_tx_factory: Callable[..., Session], delete_ledger_db: Callable[..., None], + product_manager: ProductManager, create_main_accounts: Callable[..., None], ledger_collection, task_adj_collection, @@ -799,7 +802,7 @@ class TestBusinessBalance: assert df.shape == (20, 28) gr_business.prebuild_balance( - thl_pg_config=thl_web_rr, + product_manager=product_manager, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -846,6 +849,7 @@ class TestBusinessBalance: start: datetime, bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], payout_event_manager, + product_manager: ProductManager, adj_to_fail_with_tx_factory: Callable[..., None], thl_web_rr: PostgresConfig, ledger_manager: LedgerManager, @@ -903,7 +907,7 @@ class TestBusinessBalance: 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, @@ -958,6 +962,7 @@ class TestBusinessBalance: bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], thl_ledger_manager: ThlLedgerManager, ledger_manager: LedgerManager, + product_manager: ProductManager, start: datetime, thl_web_rr: PostgresConfig, payout_event_manager, @@ -1065,7 +1070,7 @@ class TestBusinessBalance: assert df.shape == (20, 28) gr_business.prebuild_balance( - thl_pg_config=thl_web_rr, + product_manager=product_manager, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -1073,7 +1078,7 @@ class TestBusinessBalance: ) gr_business.prebuild_balance( - thl_pg_config=thl_web_rr, + product_manager=product_manager, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -1083,7 +1088,7 @@ class TestBusinessBalance: day1_bal = gr_business.balance gr_business.prebuild_balance( - thl_pg_config=thl_web_rr, + product_manager=product_manager, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -1093,7 +1098,7 @@ class TestBusinessBalance: day2_bal = gr_business.balance gr_business.prebuild_balance( - thl_pg_config=thl_web_rr, + product_manager=product_manager, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -1103,7 +1108,7 @@ class TestBusinessBalance: day3_bal = gr_business.balance gr_business.prebuild_balance( - thl_pg_config=thl_web_rr, + product_manager=product_manager, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -1113,7 +1118,7 @@ class TestBusinessBalance: day4_bal = gr_business.balance gr_business.prebuild_balance( - thl_pg_config=thl_web_rr, + product_manager=product_manager, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -1123,7 +1128,7 @@ class TestBusinessBalance: day5_bal = gr_business.balance gr_business.prebuild_balance( - thl_pg_config=thl_web_rr, + product_manager=product_manager, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, @@ -1196,8 +1201,10 @@ class TestBusinessMethods: ledger_manager: LedgerManager, thl_ledger_manager: ThlLedgerManager, business_payout_event_manager, + gr_business_bank_account_manager: BusinessBankAccountManager, + product_manager: ProductManager, product_factory: Callable[..., Product], - team: Team, + gr_team: Team, session_with_tx_factory: Callable[..., Session], user_factory: Callable[..., User], ledger_collection, @@ -1211,7 +1218,7 @@ class TestBusinessMethods: 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) + p1 = product_factory(team=gr_team, business=gr_business) u1 = user_factory(product=p1) # Business needs tx & incite to build balance @@ -1223,7 +1230,9 @@ class TestBusinessMethods: pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) gr_business.set_cache( - pg_config=gr_db, + product_manager=product_manager, + business_bank_account_manager=gr_business_bank_account_manager, + pg_config=thl_web_rr, thl_web_rr=thl_web_rr, redis_config=gr_redis_config, client=client_no_amm, @@ -1260,6 +1269,8 @@ class TestBusinessMethods: ledger_manager: LedgerManager, thl_ledger_manager: ThlLedgerManager, business_payout_event_manager, + product_manager: ProductManager, + gr_business_bank_account_manager: BusinessBankAccountManager, user_factory: Callable[..., User], delete_ledger_db: Callable[..., None], create_main_accounts: Callable[..., None], @@ -1286,6 +1297,8 @@ class TestBusinessMethods: pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) gr_business.set_cache( + product_manager=product_manager, + business_bank_account_manager=gr_business_bank_account_manager, pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config, @@ -1351,6 +1364,7 @@ class TestBusinessMethods: enriched_session_merge, client_no_amm: DaskClient, wall_collection: WallDFCollection, + product_manager: ProductManager, session_collection: SessionDFCollection, thl_web_rr: PostgresConfig, user_factory: Callable[..., User], @@ -1389,7 +1403,7 @@ class TestBusinessMethods: ) gr_business.prebuild_enriched_session_parquet( - thl_pg_config=thl_web_rr, + product_manager=product_manager, ds=mnt_filepath, client=client_no_amm, mnt_gr_api=mnt_gr_api_dir, @@ -1409,6 +1423,7 @@ class TestBusinessMethods: enriched_wall_merge, client_no_amm: DaskClient, wall_collection: WallDFCollection, + product_manager: ProductManager, session_collection: SessionDFCollection, thl_web_rr: PostgresConfig, user_factory: Callable[..., User], @@ -1447,7 +1462,7 @@ class TestBusinessMethods: ) gr_business.prebuild_enriched_wall_parquet( - thl_pg_config=thl_web_rr, + product_manager=product_manager, 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 aa2de45..a94d53f 100644 --- a/tests/models/gr/test_team.py +++ b/tests/models/gr/test_team.py @@ -29,7 +29,10 @@ if TYPE_CHECKING: from generalresearch.incite.mergers.foundations.enriched_wall import ( EnrichedWallMerge, ) + from generalresearch.managers.gr.authentication import GRUserManager + from generalresearch.managers.gr.business import BusinessManager from generalresearch.managers.gr.team import MembershipManager, TeamManager + from generalresearch.managers.thl.product import ProductManager from generalresearch.models.gr.authentication import GRUser from generalresearch.models.gr.team import Membership from generalresearch.models.thl.session import Session @@ -40,118 +43,118 @@ if TYPE_CHECKING: class TestTeam: - def test_init(self, team: Team): + def test_init(self, gr_team: Team): - assert isinstance(team, Team) - assert isinstance(team.id, int) - assert isinstance(team.uuid, str) + assert isinstance(gr_team, Team) + assert isinstance(gr_team.id, int) + assert isinstance(gr_team.uuid, str) - def test_memberships_none(self, team: Team, gr_db: PostgresConfig): - assert team.memberships is None + def test_memberships_none( + self, gr_team: Team, gr_membership_manager: MembershipManager + ): + assert gr_team.memberships is None - team.prefetch_memberships(pg_config=gr_db) - assert isinstance(team.memberships, list) - assert len(team.memberships) == 0 + gr_team.prefetch_memberships(membership_manager=gr_membership_manager) + assert isinstance(gr_team.memberships, list) + assert len(gr_team.memberships) == 0 def test_memberships( self, - team: Team, + gr_team: Team, gr_user: GRUser, gr_user_factory: Callable[..., GRUser], - membership_manager: MembershipManager, - gr_db: PostgresConfig, + gr_membership_manager: MembershipManager, ): - assert team.memberships is None + assert gr_team.memberships is None - team.prefetch_memberships(pg_config=gr_db) - assert isinstance(team.memberships, list) - assert len(team.memberships) == 1 - assert team.memberships[0].user_id == gr_user.id + gr_team.prefetch_memberships(membership_manager=gr_membership_manager) + assert isinstance(gr_team.memberships, list) + assert len(gr_team.memberships) == 1 + assert gr_team.memberships[0].user_id == gr_user.id # Create another new Membership - membership_manager.create(team=team, gr_user=gr_user_factory()) - assert len(team.memberships) == 1 - team.prefetch_memberships(pg_config=gr_db) - assert len(team.memberships) == 2 + gr_membership_manager.create(team=gr_team, gr_user=gr_user_factory()) + assert len(gr_team.memberships) == 1 + gr_team.prefetch_memberships(membership_manager=gr_membership_manager) + assert len(gr_team.memberships) == 2 def test_gr_users( self, - team: Team, + gr_team: Team, gr_user_factory: Callable[..., GRUser], membership_manager: MembershipManager, - gr_db: PostgresConfig, - gr_redis_config: RedisConfig, + gr_user_manager: GRUserManager, ): - assert team.gr_users is None + assert gr_team.gr_users is None - team.prefetch_gr_users(pg_config=gr_db, redis_config=gr_redis_config) - assert isinstance(team.gr_users, list) - assert len(team.gr_users) == 0 + gr_team.prefetch_gr_users(gr_user_manager=gr_user_manager) + assert isinstance(gr_team.gr_users, list) + assert len(gr_team.gr_users) == 0 # Create a new Membership - membership_manager.create(team=team, gr_user=gr_user_factory()) - assert len(team.gr_users) == 0 - team.prefetch_gr_users(pg_config=gr_db, redis_config=gr_redis_config) - assert len(team.gr_users) == 1 + membership_manager.create(team=gr_team, gr_user=gr_user_factory()) + assert len(gr_team.gr_users) == 0 + gr_team.prefetch_gr_users(gr_user_manager=gr_user_manager) + assert len(gr_team.gr_users) == 1 # Create another Membership - membership_manager.create(team=team, gr_user=gr_user_factory()) - assert len(team.gr_users) == 1 - team.prefetch_gr_users(pg_config=gr_db, redis_config=gr_redis_config) - assert len(team.gr_users) == 2 + membership_manager.create(team=gr_team, gr_user=gr_user_factory()) + assert len(gr_team.gr_users) == 1 + gr_team.prefetch_gr_users(gr_user_manager=gr_user_manager) + assert len(gr_team.gr_users) == 2 def test_businesses( self, - team: Team, + gr_team: Team, business: Business, team_manager: TeamManager, - gr_db: PostgresConfig, - gr_redis_config: RedisConfig, + gr_business_manager: BusinessManager, ): - assert team.businesses is None + assert gr_team.businesses is None - team.prefetch_businesses(pg_config=gr_db, redis_config=gr_redis_config) - assert isinstance(team.businesses, list) - assert len(team.businesses) == 0 + gr_team.prefetch_businesses(business_manager=gr_business_manager) + assert isinstance(gr_team.businesses, list) + assert len(gr_team.businesses) == 0 - team_manager.add_business(team=team, business=business) - assert len(team.businesses) == 0 - team.prefetch_businesses(pg_config=gr_db, redis_config=gr_redis_config) - assert len(team.businesses) == 1 - assert isinstance(team.businesses[0], Business) - assert team.businesses[0].uuid == business.uuid + team_manager.add_business(team=gr_team, business=business) + assert len(gr_team.businesses) == 0 + gr_team.prefetch_businesses(business_manager=gr_business_manager) + assert len(gr_team.businesses) == 1 + assert isinstance(gr_team.businesses[0], Business) + assert gr_team.businesses[0].uuid == business.uuid def test_products( self, - team: Team, + gr_team: Team, product_factory: Callable[..., Product], thl_web_rr: PostgresConfig, + product_manager: ProductManager, ): - assert team.products is None + assert gr_team.products is None - team.prefetch_products(thl_pg_config=thl_web_rr) - assert isinstance(team.products, list) - assert len(team.products) == 0 + gr_team.prefetch_products(product_manager=product_manager) + assert isinstance(gr_team.products, list) + assert len(gr_team.products) == 0 - product_factory(team=team) - assert len(team.products) == 0 - team.prefetch_products(thl_pg_config=thl_web_rr) - assert len(team.products) == 1 - assert isinstance(team.products[0], Product) + product_factory(team=gr_team) + assert len(gr_team.products) == 0 + gr_team.prefetch_products(product_manager=product_manager) + assert len(gr_team.products) == 1 + assert isinstance(gr_team.products[0], Product) class TestTeamMethods: - def test_cache_key(self, team: Team): - assert isinstance(team.cache_key, str) - assert ":" in team.cache_key - assert str(team.uuid) in team.cache_key + def test_cache_key(self, gr_team: Team): + assert isinstance(gr_team.cache_key, str) + assert ":" in gr_team.cache_key + assert str(gr_team.uuid) in gr_team.cache_key def test_set_cache( self, - team: Team, + gr_team: Team, gr_db: PostgresConfig, thl_web_rr: PostgresConfig, gr_redis_config: RedisConfig, @@ -162,9 +165,9 @@ class TestTeamMethods: enriched_session_merge: EnrichedSessionMerge, ): client = gr_redis_config.create_redis_client() - assert client.get(name=team.cache_key) is None + assert client.get(name=gr_team.cache_key) is None - team.set_cache( + gr_team.set_cache( pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config, @@ -175,7 +178,7 @@ class TestTeamMethods: enriched_session=enriched_session_merge, ) - assert client.hgetall(name=team.cache_key) is not None + assert client.hgetall(name=gr_team.cache_key) is not None def test_set_cache_team( self, @@ -183,7 +186,7 @@ class TestTeamMethods: gr_db: PostgresConfig, thl_web_rr: PostgresConfig, product_factory: Callable[..., Product], - team: Team, + gr_team: Team, membership_factory: Callable[..., Membership], gr_redis_config: RedisConfig, mnt_filepath: GRLDatasets, @@ -193,10 +196,10 @@ class TestTeamMethods: ): from generalresearch.models.gr.team import Team - p1 = product_factory(team=team) - membership_factory(team=team, gr_user=gr_user) + p1 = product_factory(team=gr_team) + membership_factory(team=gr_team, gr_user=gr_user) - team.set_cache( + gr_team.set_cache( pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config, @@ -208,7 +211,7 @@ class TestTeamMethods: ) team2 = Team.from_redis( - uuid=team.uuid, + uuid=gr_team.uuid, fields=["id", "memberships", "gr_users", "businesses", "products"], gr_redis_config=gr_redis_config, ) @@ -216,7 +219,7 @@ class TestTeamMethods: assert isinstance(team2, Team) assert isinstance(team2.products, list) assert isinstance(team2.gr_users, list) - assert team.model_dump_json() == team2.model_dump_json() + assert gr_team.model_dump_json() == team2.model_dump_json() assert p1.uuid in [p.uuid for p in team2.products] assert len(team2.gr_users) == 1 assert gr_user.id in [gru.id for gru in team2.gr_users] @@ -235,14 +238,14 @@ class TestTeamMethods: delete_df_collection: Callable[..., None], mnt_filepath: GRLDatasets, mnt_gr_api_dir: Path, - team: Team, + gr_team: Team, ): delete_df_collection(coll=wall_collection) delete_df_collection(coll=session_collection) - p1 = product_factory(team=team) - p2 = product_factory(team=team) + p1 = product_factory(team=gr_team) + p2 = product_factory(team=gr_team) for p in [p1, p2]: u = user_factory(product=p) @@ -263,7 +266,7 @@ class TestTeamMethods: pg_config=thl_web_rr, ) - team.prebuild_enriched_session_parquet( + gr_team.prebuild_enriched_session_parquet( thl_pg_config=thl_web_rr, ds=mnt_filepath, client=client_no_amm, @@ -273,7 +276,7 @@ class TestTeamMethods: # Now try to read from path df = pd.read_parquet( - os.path.join(mnt_gr_api_dir, "pop_session", f"{team.file_key}.parquet") + os.path.join(mnt_gr_api_dir, "pop_session", f"{gr_team.file_key}.parquet") ) assert isinstance(df, pd.DataFrame) @@ -291,14 +294,14 @@ class TestTeamMethods: delete_df_collection: Callable[..., None], mnt_filepath: GRLDatasets, mnt_gr_api_dir: Path, - team: Team, + gr_team: Team, ): delete_df_collection(coll=wall_collection) delete_df_collection(coll=session_collection) - p1 = product_factory(team=team) - p2 = product_factory(team=team) + p1 = product_factory(team=gr_team) + p2 = product_factory(team=gr_team) for p in [p1, p2]: u = user_factory(product=p) @@ -319,7 +322,7 @@ class TestTeamMethods: pg_config=thl_web_rr, ) - team.prebuild_enriched_wall_parquet( + gr_team.prebuild_enriched_wall_parquet( thl_pg_config=thl_web_rr, ds=mnt_filepath, client=client_no_amm, @@ -329,6 +332,6 @@ class TestTeamMethods: # Now try to read from path df = pd.read_parquet( - os.path.join(mnt_gr_api_dir, "pop_event", f"{team.file_key}.parquet") + os.path.join(mnt_gr_api_dir, "pop_event", f"{gr_team.file_key}.parquet") ) assert isinstance(df, pd.DataFrame) -- cgit v1.2.3 From 17ff15c06655717627da820417337c6b0b97de42 Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Wed, 2 Sep 2026 23:31:52 -0700 Subject: Lots more tests/managers/thl - doing all the factory organization from create_dummy --- generalresearch/models/gr/team.py | 6 +- generalresearch/thl_django/app/test_settings.py | 2 +- test_utils/grliq/conftest.py | 174 +++++++----- test_utils/managers/conftest.py | 10 - test_utils/managers/thl/conftest.py | 53 ++++ test_utils/models/conftest.py | 107 ++----- test_utils/models/thl/conftest.py | 308 ++++++++++++++++----- tests/grliq/managers/test_forensic_data.py | 65 +++-- tests/grliq/managers/test_forensic_results.py | 11 +- tests/managers/gr/test_business.py | 38 +-- tests/managers/thl/test_contest/test_milestone.py | 2 +- tests/managers/thl/test_contest/test_raffle.py | 10 +- tests/managers/thl/test_ipinfo.py | 20 +- tests/managers/thl/test_ledger/test_lm_accounts.py | 4 +- tests/managers/thl/test_ledger/test_lm_tx_locks.py | 2 +- tests/managers/thl/test_ledger/test_thl_lm_tx.py | 3 +- tests/managers/thl/test_ledger/test_wallet.py | 6 +- tests/managers/thl/test_product.py | 76 +++-- tests/managers/thl/test_task_status.py | 18 +- tests/managers/thl/test_user_manager/test_base.py | 20 +- tests/managers/thl/test_wall_manager.py | 2 +- tests/models/gr/test_business.py | 3 - tests/models/gr/test_team.py | 28 +- tests/models/thl/test_product.py | 6 +- tests/models/thl/test_user.py | 10 +- 25 files changed, 624 insertions(+), 360 deletions(-) (limited to 'test_utils/models') diff --git a/generalresearch/models/gr/team.py b/generalresearch/models/gr/team.py index aa62c5a..e10dba5 100644 --- a/generalresearch/models/gr/team.py +++ b/generalresearch/models/gr/team.py @@ -132,8 +132,8 @@ class Team(BaseModel): def prefetch_gr_users(self, gr_user_manager: GRUserManager) -> None: self.gr_users = gr_user_manager.get_by_team(team_id=self.id) - def prefetch_businesses(self, business_manager: BusinessManager) -> None: - self.businesses = business_manager.get_by_team(team_id=self.id) + def prefetch_businesses(self, gr_business_manager: BusinessManager) -> None: + self.businesses = gr_business_manager.get_by_team(team_id=self.id) def prefetch_products(self, product_manager: ProductManager) -> None: self.products = product_manager.fetch_uuids(team_uuids=[self.uuid]) @@ -273,7 +273,7 @@ class Team(BaseModel): ) -> None: self.prefetch_products(product_manager=product_manager) self.prefetch_gr_users(gr_user_manager=gr_user_manager) - self.prefetch_businesses(business_manager=gr_business_manager) + self.prefetch_businesses(gr_business_manager=gr_business_manager) self.prefetch_memberships(membership_manager=gr_membership_manager) rc = redis_config.create_redis_client() diff --git a/generalresearch/thl_django/app/test_settings.py b/generalresearch/thl_django/app/test_settings.py index f3d23af..d6ab124 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-728bcf', + "NAME": 'unittest-2026-09-03-44c0b4', "USER": 'jenkins', "PASSWORD": '123456789', "HOST": 'unittest-postgresql.fmt2.grl.internal', diff --git a/test_utils/grliq/conftest.py b/test_utils/grliq/conftest.py index bb1a167..f399a99 100644 --- a/test_utils/grliq/conftest.py +++ b/test_utils/grliq/conftest.py @@ -27,7 +27,8 @@ if TYPE_CHECKING: GrlIqEventManager, ) -# === Miscellaneous === + +# --- Assets --- @pytest.fixture(scope="function") @@ -48,18 +49,98 @@ def grliq_db(postgres_instance: PostgresDsn) -> PostgresConfig: ) -# === Managers === +# --- GRLIQ Data --- @pytest.fixture(scope="session") -def grliq_dm(grliq_db: PostgresConfig) -> GrlIqDataManager: +def grliq_data_manager(grliq_db: PostgresConfig) -> GrlIqDataManager: assert grliq_db.dsn.path assert "/unittest-" in grliq_db.dsn.path return GrlIqDataManager(postgres_config=grliq_db) @pytest.fixture(scope="session") -def grliq_em(grliq_db: PostgresConfig) -> GrlIqEventManager: +def grliq_dm(grliq_data_manager: GrlIqDataManager) -> GrlIqDataManager: + return grliq_data_manager + + +@pytest.fixture +def grliq_data_factory( + grliq_data_manager: GrlIqDataManager, grliq_data_list: list[dict[str, Any]] +) -> Callable[..., GrlIqData]: + + def _inner( + save: bool = True, + is_attempt_allowed: bool = True, + product_id: str | None = None, + product_user_id: str | None = None, + uuid: str | None = None, + mid: str | None = None, + created_at: datetime | None = None, + ) -> GrlIqData: + """ + Creates a dummy record in the db with a GrlIqData (data), GrlIqCheckerResults (result_data), + and GrlIqForensicCategoryResult (category_results) + :param is_attempt_allowed: Whether the attempt is allowed. + :param product_id: product_id of user + :param product_user_id: product_user_id of user + :param uuid: uuid for the grliq data record + :param mid: the thl_session:uuid / mid for the attempt. + :return: + """ + + if save: + res: GrlIqData = grliq_data_list[int(is_attempt_allowed)]["data"] + + product_id = product_id or uuid4().hex + 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=UTC) + + res["data"].product_id = product_id + res["data"].product_user_id = product_user_id + res["data"].uuid = uuid + res["data"].mid = mid + res["data"].created_at = created_at + res["result_data"].uuid = uuid + res["category_result"].uuid = uuid + + return grliq_data_manager.create( + iq_data=res["data"], + result_data=res["result_data"], + category_result=res["category_result"], + fraud_score=res["category_result"].fraud_score, + is_attempt_allowed=res["category_result"].is_attempt_allowed(), + ) + else: + raise ValueError("Unsaved GRLIQ Data not supported yet") + + return _inner + + +@pytest.fixture(scope="function") +def grliq_data(grliq_data_list: list[dict[str, Any]]) -> GrlIqData: + + g: GrlIqData = grliq_data_list[1]["data"] + + g.id = None + g.uuid = uuid4().hex + g.created_at = datetime.now(tz=UTC) + g.timestamp = g.created_at - timedelta(seconds=10) + return g + + +@pytest.fixture(scope="function") +def unsaved_grliq_data(grliq_data_list: list[dict[str, Any]]) -> GrlIqData: + raise ValueError("Not supported") + + +# --- GRLIQ Event --- + + +@pytest.fixture(scope="session") +def grliq_event_manager(grliq_db: PostgresConfig) -> GrlIqEventManager: assert grliq_db.dsn.path assert "/unittest-" in grliq_db.dsn.path @@ -71,16 +152,36 @@ def grliq_em(grliq_db: PostgresConfig) -> GrlIqEventManager: @pytest.fixture(scope="session") -def grliq_crr(grliq_db: PostgresConfig) -> GrlIqCategoryResultsReader: +def grliq_em(grliq_event_manager: GrlIqEventManager) -> GrlIqEventManager: + return grliq_event_manager + + +# --- GRLIQ Category Results Reader --- + + +@pytest.fixture(scope="session") +def grliq_category_results_reader( + grliq_db: PostgresConfig, +) -> GrlIqCategoryResultsReader: assert grliq_db.dsn.path assert "/unittest-" in grliq_db.dsn.path return GrlIqCategoryResultsReader(postgres_config=grliq_db) +@pytest.fixture(scope="session") +def grliq_crr( + grliq_category_results_reader: GrlIqCategoryResultsReader, +) -> GrlIqCategoryResultsReader: + return grliq_category_results_reader + + # === Models === +# === Miscellaneous === + + @pytest.fixture(scope="session") def grliq_data_list() -> list[dict[str, Any]]: return [ @@ -111,66 +212,3 @@ def grliq_data_list() -> list[dict[str, Any]]: "is_attempt_allowed": True, }, ] - - -@pytest.fixture(scope="function") -def grliq_data(grliq_data_list: list[dict[str, Any]]) -> GrlIqData: - - g: GrlIqData = grliq_data_list[1]["data"] - - g.id = None - g.uuid = uuid4().hex - g.created_at = datetime.now(tz=UTC) - g.timestamp = g.created_at - timedelta(seconds=10) - return g - - -@pytest.fixture -def grliq_data_factory( - grliq_dm: GrlIqDataManager, grliq_data_list: list[dict[str, Any]] -) -> Callable[..., GrlIqData]: - - def _inner( - is_attempt_allowed: bool = True, - product_id: str | None = None, - product_user_id: str | None = None, - uuid: str | None = None, - mid: str | None = None, - created_at: datetime | None = None, - ) -> GrlIqData: - """ - Creates a dummy record in the db with a GrlIqData (data), GrlIqCheckerResults (result_data), - and GrlIqForensicCategoryResult (category_results) - :param is_attempt_allowed: Whether the attempt is allowed. - :param product_id: product_id of user - :param product_user_id: product_user_id of user - :param uuid: uuid for the grliq data record - :param mid: the thl_session:uuid / mid for the attempt. - :return: - """ - - res: GrlIqData = grliq_data_list[int(is_attempt_allowed)]["data"] - - product_id = product_id or uuid4().hex - 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=UTC) - - res["data"].product_id = product_id - res["data"].product_user_id = product_user_id - res["data"].uuid = uuid - res["data"].mid = mid - res["data"].created_at = created_at - res["result_data"].uuid = uuid - res["category_result"].uuid = uuid - - return grliq_dm.create( - iq_data=res["data"], - result_data=res["result_data"], - category_result=res["category_result"], - fraud_score=res["category_result"].fraud_score, - is_attempt_allowed=res["category_result"].is_attempt_allowed(), - ) - - return _inner diff --git a/test_utils/managers/conftest.py b/test_utils/managers/conftest.py index ff088c2..3e7b304 100644 --- a/test_utils/managers/conftest.py +++ b/test_utils/managers/conftest.py @@ -55,16 +55,6 @@ def ip_geoname_manager(thl_web_rw: PostgresConfig) -> IPGeonameManager: return IPGeonameManager(pg_config=thl_web_rw) -@pytest.fixture(scope="session") -def ip_information_manager(thl_web_rw: PostgresConfig) -> IPInformationManager: - assert thl_web_rw.dsn.path - assert "/unittest-" in thl_web_rw.dsn.path - - from generalresearch.managers.thl.ipinfo import IPInformationManager - - return IPInformationManager(pg_config=thl_web_rw) - - @pytest.fixture(scope="session") def ip_record_manager( thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig diff --git a/test_utils/managers/thl/conftest.py b/test_utils/managers/thl/conftest.py index 18a31e2..8ca4383 100644 --- a/test_utils/managers/thl/conftest.py +++ b/test_utils/managers/thl/conftest.py @@ -23,6 +23,10 @@ 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.ipinfo import ( + IPGeonameManager, + IPInformationManager, + ) from generalresearch.managers.thl.payout import ( BrokerageProductPayoutEventManager, BusinessPayoutEventManager, @@ -40,6 +44,10 @@ if TYPE_CHECKING: from generalresearch.managers.thl.user_manager.user_metadata_manager import ( UserMetadataManager, ) + from generalresearch.managers.thl.userhealth import ( + AuditLogManager, + IPRecordManager, + ) from generalresearch.managers.thl.wall import ( WallCacheManager, WallManager, @@ -153,6 +161,13 @@ def brokerage_product_payout_event_manager( ) +@pytest.fixture() +def audit_log_manager(thl_web_rw: PostgresConfig) -> AuditLogManager: + from generalresearch.managers.thl.userhealth import AuditLogManager + + return AuditLogManager(pg_config=thl_web_rw) + + @pytest.fixture(scope="session") def business_payout_event_manager( thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig @@ -319,3 +334,41 @@ def surveypenalty_manager(thl_redis_config: RedisConfig): from generalresearch.managers.thl.survey_penalty import SurveyPenaltyManager return SurveyPenaltyManager(redis_config=thl_redis_config) + + +# --- IP Geolocation --- + + +@pytest.fixture +def ip_geoname_manager(thl_web_rw: PostgresConfig) -> IPGeonameManager: + from generalresearch.managers.thl.ipinfo import IPGeonameManager + + return IPGeonameManager(pg_config=thl_web_rw) + + +# --- IP Information --- + + +@pytest.fixture(scope="session") +def ip_information_manager(thl_web_rw: PostgresConfig) -> IPInformationManager: + assert thl_web_rw.dsn.path + assert "/unittest-" in thl_web_rw.dsn.path + + from generalresearch.managers.thl.ipinfo import IPInformationManager + + return IPInformationManager(pg_config=thl_web_rw) + + +# --- IP Record --- + + +@pytest.fixture(scope="session") +def ip_record_manager( + thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig +) -> IPRecordManager: + assert thl_web_rw.dsn.path + assert "/unittest-" in thl_web_rw.dsn.path + + from generalresearch.managers.thl.userhealth import IPRecordManager + + return IPRecordManager(pg_config=thl_web_rw, redis_config=thl_redis_config) diff --git a/test_utils/models/conftest.py b/test_utils/models/conftest.py index d71593f..d5c9a71 100644 --- a/test_utils/models/conftest.py +++ b/test_utils/models/conftest.py @@ -10,6 +10,7 @@ from uuid import uuid4 import pytest from pydantic import AwareDatetime, PositiveInt +from pytest import FixtureRequest from pytest import FixtureRequest as Request from generalresearch.models.definitions import Source @@ -50,8 +51,6 @@ if TYPE_CHECKING: ) from generalresearch.models.thl.session import Session, 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.pg_helper import PostgresConfig # === THL === @@ -59,15 +58,15 @@ if TYPE_CHECKING: @pytest.fixture def user( - request, - product_manager: ProductManager, + 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_manager.create_dummy() + product = product_factory() u = user_manager.create_dummy(product_id=product.id) u.prefetch_product(pg_config=thl_web_rr) @@ -309,31 +308,35 @@ def payout_config(request: Request) -> PayoutConfig: @pytest.fixture def product_user_wallet_yes( - payout_config: PayoutConfig, product_manager: ProductManager + product_factory: Callable[..., Product], + payout_config: PayoutConfig, + product_manager: ProductManager, ) -> Product: from generalresearch.models.thl.product import UserWalletConfig - return product_manager.create_dummy( + return product_factory( payout_config=payout_config, user_wallet_config=UserWalletConfig(enabled=True) ) @pytest.fixture -def product_user_wallet_no(product_manager: ProductManager) -> Product: +def product_user_wallet_no( + product_factory: Callable[..., Product], product_manager: ProductManager +) -> Product: from generalresearch.models.thl.product import UserWalletConfig - return product_manager.create_dummy( - user_wallet_config=UserWalletConfig(enabled=False) - ) + return product_factory(user_wallet_config=UserWalletConfig(enabled=False)) @pytest.fixture def product_amt_true( - product_manager: ProductManager, payout_config: PayoutConfig + product_factory: Callable[..., Product], + product_manager: ProductManager, + payout_config: PayoutConfig, ) -> Product: from generalresearch.models.thl.product import UserWalletConfig - return product_manager.create_dummy( + return product_factory( user_wallet_config=UserWalletConfig(amt=True, enabled=True), payout_config=payout_config, ) @@ -370,84 +373,6 @@ def bp_payout_factory( return _inner -@pytest.fixture -def audit_log(audit_log_manager: AuditLogManager, user: User) -> AuditLog: - - return audit_log_manager.create_dummy(user_id=user.user_id) - - -@pytest.fixture -def audit_log_factory( - audit_log_manager: AuditLogManager, -) -> Callable[..., AuditLog]: - - def _inner( - user_id: PositiveInt, - level: AuditLogLevel | None = None, - event_type: str | None = None, - event_msg: str | None = None, - event_value: float | None = None, - ) -> AuditLog: - return audit_log_manager.create_dummy( - user_id=user_id, - level=level, - event_type=event_type, - event_msg=event_msg, - event_value=event_value, - ) - - return _inner - - -@pytest.fixture -def ip_geoname(ip_geoname_manager: IPGeonameManager) -> IPGeoname: - return ip_geoname_manager.create_dummy() - - -@pytest.fixture -def ip_information( - ip_information_manager: IPInformationManager, ip_geoname: IPGeoname -) -> IPInformation: - return ip_information_manager.create_dummy( - geoname_id=ip_geoname.geoname_id, country_iso=ip_geoname.country_iso - ) - - -@pytest.fixture -def ip_information_factory( - ip_information_manager: IPInformationManager, -) -> Callable[..., IPInformation]: - - def _inner(ip: str, geoname: IPGeoname, **kwargs) -> IPInformation: - return ip_information_manager.create_dummy( - ip=ip, - geoname_id=geoname.geoname_id, - country_iso=geoname.country_iso, - **kwargs, - ) - - return _inner - - -@pytest.fixture -def ip_record( - ip_record_manager: IPRecordManager, ip_geoname: IPGeoname, user: User -) -> IPRecord: - - return ip_record_manager.create_dummy(user_id=user.user_id) - - -@pytest.fixture -def ip_record_factory( - ip_record_manager: IPRecordManager, user: User -) -> Callable[..., IPRecord]: - - def _inner(user_id: PositiveInt, ip: str | None = None) -> IPRecord: - return ip_record_manager.create_dummy(user_id=user_id, ip=ip) - - return _inner - - @pytest.fixture(scope="session") def buyer(buyer_manager: BuyerManager) -> Buyer: buyer_code = uuid4().hex diff --git a/test_utils/models/thl/conftest.py b/test_utils/models/thl/conftest.py index 5826f0d..14f8f36 100644 --- a/test_utils/models/thl/conftest.py +++ b/test_utils/models/thl/conftest.py @@ -62,6 +62,7 @@ if TYPE_CHECKING: from generalresearch.models.thl.user_iphistory import IPRecord from generalresearch.models.thl.userhealth import AuditLog from generalresearch.models.thl.wallet.cashout_method import CashMailOrderData + from generalresearch.pg_helper import PostgresConfig fake = faker.Faker() @@ -71,30 +72,6 @@ def wall_status() -> Status: return Status.COMPLETE -@pytest.fixture -def user_factory(user_manager: UserManager) -> Callable[..., User]: - - def _inner( - # --- Create dummy "optional" --- # - product_user_id: str | None = None, - # --- Optional --- # - product_id: UUIDStr | None = None, - product: Product | None = None, - created: datetime | None = None, - ) -> User: - - product_user_id = product_user_id or uuid4().hex - - return user_manager.create_user( - product_user_id=product_user_id, - product_id=product_id, - product=product, - created=created, - ) - - return _inner - - @pytest.fixture def wall_factory( wall_manager: WallManager, session_factory: Session @@ -155,9 +132,10 @@ def product_factory(product_manager: ProductManager) -> Callable[..., Product]: def _inner( save: bool = True, team: Team | None = None, + team_id: UUIDStr | None = None, business: Business | None = None, - product_id: UUIDStr | None = None, business_id: UUIDStr | None = None, + product_id: UUIDStr | None = None, name: str | None = None, redirect_url: str | None = None, harmonizer_domain: str | None = None, @@ -174,8 +152,10 @@ def product_factory(product_manager: ProductManager) -> Callable[..., Product]: product_id = product_id if product_id else uuid4().hex - team_id = team.uuid if team else uuid4().hex - business_id = business.uuid if business else uuid4().hex + team_id = (team.uuid if team else None) or team_id or uuid4().hex + business_id = ( + (business.uuid if business else None) or business_id or uuid4().hex + ) name = name if name else f"name-{product_id[:12]}" redirect_url = redirect_url if redirect_url else "https://www.example.com/" @@ -256,9 +236,12 @@ def session_factory(session_manager: SessionManager): @pytest.fixture -def ipgeoname_factory(ipgeoname_manager: IPGeonameManager) -> Callable[..., IPGeoname]: +def ip_geoname_factory( + ip_geoname_manager: IPGeonameManager, +) -> Callable[..., IPGeoname]: def _inner( + save: bool, geoname_id: PositiveInt | None = None, continent_code: str | None = None, continent_name: str | None = None, @@ -273,31 +256,47 @@ def ipgeoname_factory(ipgeoname_manager: IPGeonameManager) -> Callable[..., IPGe time_zone: str | None = None, is_in_european_union: bool | None = None, ) -> IPGeoname: - - return ipgeoname_manager.create( - geoname_id=geoname_id or randint(1, 999_999_999), - continent_code=continent_code or "na", - continent_name=continent_name or "North America", - country_iso=country_iso or "us", - country_name=country_name or "United States", - subdivision_1_iso=subdivision_1_iso or "fl", - subdivision_1_name=subdivision_1_name or "Florida", - subdivision_2_iso=subdivision_2_iso, - subdivision_2_name=subdivision_2_name, - city_name=city_name, - metro_code=metro_code, - time_zone=time_zone, - is_in_european_union=is_in_european_union, - ) + if save: + return ip_geoname_manager.create( + geoname_id=geoname_id or randint(1, 999_999_999), + continent_code=continent_code or "na", + continent_name=continent_name or "North America", + country_iso=country_iso or "us", + country_name=country_name or "United States", + subdivision_1_iso=subdivision_1_iso or "fl", + subdivision_1_name=subdivision_1_name or "Florida", + subdivision_2_iso=subdivision_2_iso, + subdivision_2_name=subdivision_2_name, + city_name=city_name, + metro_code=metro_code, + time_zone=time_zone, + is_in_european_union=is_in_european_union, + ) + else: + raise ValueError("Unsaved IPGeoname not yet supported") return _inner -def ipinformation_factory( +@pytest.fixture() +def ip_geoname(ip_geoname_factory: Callable[..., IPGeoname]) -> IPGeoname: + return ip_geoname_factory(save=True) + + +@pytest.fixture() +def unsaved_ip_geoname(ip_geoname_factory: Callable[..., IPGeoname]) -> IPGeoname: + return ip_geoname_factory(save=True) + + +# --- IP Information --- + + +def ip_information_factory( ipinformation_manager: IPInformationManager, ) -> Callable[..., IPInformation]: def _inner( + save: bool = True, ip: IPvAnyAddressStr | None = None, geoname_id: PositiveInt | None = None, country_iso: str | None = None, @@ -324,36 +323,184 @@ def ipinformation_factory( accuracy_radius: int | None = None, ) -> IPInformation: - return ipinformation_manager.create( - ip=ip or fake.ipv4_public(), - geoname_id=geoname_id, - country_iso=country_iso or fake.country_code(), - registered_country_iso=registered_country_iso, - is_anonymous=is_anonymous, - is_anonymous_vpn=is_anonymous_vpn, - is_hosting_provider=is_hosting_provider, - is_public_proxy=is_public_proxy, - is_tor_exit_node=is_tor_exit_node, - is_residential_proxy=is_residential_proxy, - autonomous_system_number=autonomous_system_number, - autonomous_system_organization=autonomous_system_organization, - domain=domain, - isp=isp, - mobile_country_code=mobile_country_code, - mobile_network_code=mobile_network_code, - network=network, - organization=organization, - static_ip_score=static_ip_score, - user_type=user_type, - postal_code=postal_code, - latitude=latitude, - longitude=longitude, - accuracy_radius=accuracy_radius, - ) + if save: + return ipinformation_manager.create( + ip=ip or fake.ipv4_public(), + geoname_id=geoname_id, + country_iso=country_iso or fake.country_code(), + registered_country_iso=registered_country_iso, + is_anonymous=is_anonymous, + is_anonymous_vpn=is_anonymous_vpn, + is_hosting_provider=is_hosting_provider, + is_public_proxy=is_public_proxy, + is_tor_exit_node=is_tor_exit_node, + is_residential_proxy=is_residential_proxy, + autonomous_system_number=autonomous_system_number, + autonomous_system_organization=autonomous_system_organization, + domain=domain, + isp=isp, + mobile_country_code=mobile_country_code, + mobile_network_code=mobile_network_code, + network=network, + organization=organization, + static_ip_score=static_ip_score, + user_type=user_type, + postal_code=postal_code, + latitude=latitude, + longitude=longitude, + accuracy_radius=accuracy_radius, + ) + else: + raise ValueError("Unsaved IP Information not supported yet") + + return _inner + + +@pytest.fixture +def ip_information( + ip_information_factory: Callable[..., IPInformation], +) -> IPInformation: + return ip_information_factory(save=True) + + +@pytest.fixture +def unsaved_ip_information( + ip_information_factory: Callable[..., IPInformation], +) -> IPInformation: + return ip_information_factory(save=False) + + +# --- 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, + # ) + + def _inner( + user_id: PositiveInt, save: bool = True, ip: str | None = None + ) -> IPRecord: + if save: + return ip_record_manager.create_dummy(user_id=user_id, ip=ip) + else: + raise ValueError("Unsaved IP Record not supported") return _inner +@pytest.fixture() +def ip_record( + ip_record_manager: IPRecordManager, ip_geoname: IPGeoname, user: User +) -> IPRecord: + return ip_record_factory(save=True) + + +@pytest.fixture() +def unsaved_ip_record(ip_record_factory: Callable[..., IPRecord]) -> IPRecord: + return ip_record_factory(save=False) + + +# --- User --- + + +@pytest.fixture() +def user_factory( + user_manager: UserManager, thl_web_rr: PostgresConfig +) -> Callable[..., User]: + + def _inner( + save: bool = True, + # --- Create dummy "optional" --- # + product_user_id: str | None = None, + # --- Optional --- # + product_id: UUIDStr | None = None, + product: Product | None = None, + created: datetime | None = None, + ) -> User: + if save: + if product is None: + product = product_factory() + + product_user_id = product_user_id or uuid4().hex + + u = user_manager.create_user( + product_user_id=product_user_id, + product_id=product_id, + product=product, + created=created, + ) + + u = user_manager.create_dummy(product=product, created=created) + + u.prefetch_product(pg_config=thl_web_rr) + return u + + else: + raise ValueError("Unsaved User not supported") + + return _inner + + +@pytest.fixture() +def user( + user_factory: Callable[..., User], +) -> User: + return user_factory(save=True) + + +@pytest.fixture() +def unsaved_user( + user_factory: Callable[..., User], +) -> User: + return user_factory(save=False) + + +@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(save=True, 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(save=True, product=product_amt_true) + + +# --- User Payout --- + + @pytest.fixture def user_payout_event_factory( user_payout_event_manager: UserPayoutEventManager, @@ -437,11 +584,11 @@ def iprecord_factory(iprecord_manager: IPRecordManager) -> Callable[..., IPRecor return _inner -# class AuditLogManager(PostgresManager): +# --- Audit Log Manager --- -@pytest.fixture -def auditlog_factory(audit_log_manager: AuditLogManager): +@pytest.fixture() +def audit_log_factory(audit_log_manager: AuditLogManager) -> Callable[..., AuditLog]: def _inner( user_id: PositiveInt, @@ -468,6 +615,19 @@ def auditlog_factory(audit_log_manager: AuditLogManager): return _inner +@pytest.fixture() +def audit_log(auditlog_factory: Callable[..., AuditLog]) -> AuditLog: + return auditlog_factory(save=True) + + +@pytest.fixture() +def unsaved_audit_log(auditlog_factory: Callable[..., AuditLog]) -> AuditLog: + return auditlog_factory(save=False) + + +# --- --- + + @pytest.fixture(scope="session") def profiling_info_json() -> str: return ( diff --git a/tests/grliq/managers/test_forensic_data.py b/tests/grliq/managers/test_forensic_data.py index e4854e8..1b83757 100644 --- a/tests/grliq/managers/test_forensic_data.py +++ b/tests/grliq/managers/test_forensic_data.py @@ -1,5 +1,6 @@ from __future__ import annotations +from collections.abc import Callable from datetime import timedelta from typing import TYPE_CHECKING from uuid import uuid4 @@ -16,6 +17,8 @@ from generalresearch.grliq.models.forensic_result import ( if TYPE_CHECKING: from generalresearch.grliq.managers.forensic_data import ( GrlIqDataManager, + ) + from generalresearch.grliq.managers.forensic_events import ( GrlIqEventManager, ) from generalresearch.models.thl.product import Product @@ -28,10 +31,13 @@ except ImportError: class TestGrlIqDataManager: - def test_create_dummy(self, grliq_dm: GrlIqDataManager): + def test_create_dummy( + self, + grliq_data_factory: Callable[..., GrlIqData], + ): from generalresearch.grliq.models.forensic_data import GrlIqData - gd1: GrlIqData = grliq_dm.create_dummy(is_attempt_allowed=True) + gd1: GrlIqData = grliq_data_factory(is_attempt_allowed=True) assert isinstance(gd1, GrlIqData) assert isinstance(gd1.results, GrlIqCheckerResults) @@ -119,7 +125,9 @@ class TestGrlIqDataManager: class TestForensicDataGetAndFilter: - def test_events(self, grliq_dm: GrlIqDataManager): + def test_events( + self, grliq_dm: GrlIqDataManager, grliq_data_factory: Callable[..., GrlIqData] + ): """If load_events=True, the events and mouse_events attributes should be an array no matter what. An empty array means that the events were loaded, but there were no events available. @@ -129,7 +137,7 @@ class TestForensicDataGetAndFilter: """ # Load Events == False forensic_uuid = uuid4().hex - grliq_dm.create_dummy(is_attempt_allowed=True, uuid=forensic_uuid) + grliq_data_factory(is_attempt_allowed=True, uuid=forensic_uuid) instance = grliq_dm.filter_data(uuids=[forensic_uuid])[0] assert isinstance(instance, GrlIqData) @@ -144,41 +152,53 @@ class TestForensicDataGetAndFilter: assert len(instance.events) == 0 assert len(instance.mouse_events) == 0 - def test_timing(self, grliq_dm: GrlIqDataManager, grliq_em: GrlIqEventManager): + def test_timing( + self, + grliq_data_factory: Callable[..., GrlIqData], + grliq_data_manager: GrlIqDataManager, + grliq_event_manager: GrlIqEventManager, + ): forensic_uuid = uuid4().hex - grliq_dm.create_dummy(is_attempt_allowed=True, uuid=forensic_uuid) + grliq_data_factory(is_attempt_allowed=True, uuid=forensic_uuid) - instance = grliq_dm.filter_data(uuids=[forensic_uuid])[0] + instance = grliq_data_manager.filter_data(uuids=[forensic_uuid])[0] - grliq_em.update_or_create_timing( + grliq_event_manager.update_or_create_timing( session_uuid=instance.mid, timing_data=TimingData( client_rtts=[100, 200, 150], server_rtts=[150, 120, 120] ), ) - instance = grliq_dm.get_data(forensic_uuid=forensic_uuid, load_events=True) + instance = grliq_data_manager.get_data( + forensic_uuid=forensic_uuid, load_events=True + ) assert isinstance(instance, GrlIqData) assert isinstance(instance.events, list) assert isinstance(instance.mouse_events, list) assert isinstance(instance.timing_data, TimingData) def test_events_events( - self, grliq_dm: GrlIqDataManager, grliq_em: GrlIqEventManager + self, + grliq_data_factory: Callable[..., GrlIqData], + grliq_data_manager: GrlIqDataManager, + grliq_event_manager: GrlIqEventManager, ): forensic_uuid = uuid4().hex - grliq_dm.create_dummy(is_attempt_allowed=True, uuid=forensic_uuid) + grliq_data_factory(is_attempt_allowed=True, uuid=forensic_uuid) - instance = grliq_dm.filter_data(uuids=[forensic_uuid])[0] + instance = grliq_data_manager.filter_data(uuids=[forensic_uuid])[0] - grliq_em.update_or_create_events( + grliq_event_manager.update_or_create_events( session_uuid=instance.mid, events=[{"a": "b"}], mouse_events=[], event_start=instance.created_at, event_end=instance.created_at + timedelta(minutes=1), ) - instance = grliq_dm.get_data(forensic_uuid=forensic_uuid, load_events=True) + instance = grliq_data_manager.get_data( + forensic_uuid=forensic_uuid, load_events=True + ) assert isinstance(instance, GrlIqData) assert isinstance(instance.events, list) assert isinstance(instance.mouse_events, list) @@ -189,11 +209,16 @@ class TestForensicDataGetAndFilter: assert len(instance.keyboard_events) == 0 def test_events_click( - self, grliq_dm: GrlIqDataManager, grliq_em: GrlIqEventManager + self, + grliq_data_factory: Callable[..., GrlIqData], + grliq_data_manager: GrlIqDataManager, + grliq_event_manager: GrlIqEventManager, ): forensic_uuid = uuid4().hex - grliq_dm.create_dummy(is_attempt_allowed=True, uuid=forensic_uuid) - instance = grliq_dm.get_data(forensic_uuid=forensic_uuid, load_events=True) + grliq_data_factory(is_attempt_allowed=True, uuid=forensic_uuid) + instance = grliq_data_manager.get_data( + forensic_uuid=forensic_uuid, load_events=True + ) click_event = { "type": "click", @@ -203,14 +228,16 @@ class TestForensicDataGetAndFilter: "pointerType": "mouse", } me = MouseEvent.from_dict(click_event) - grliq_em.update_or_create_events( + grliq_event_manager.update_or_create_events( session_uuid=instance.mid, events=[click_event], mouse_events=[], event_start=instance.created_at, event_end=instance.created_at + timedelta(minutes=1), ) - instance = grliq_dm.get_data(forensic_uuid=forensic_uuid, load_events=True) + instance = grliq_data_manager.get_data( + forensic_uuid=forensic_uuid, load_events=True + ) assert isinstance(instance, GrlIqData) assert isinstance(instance.events, list) assert isinstance(instance.mouse_events, list) diff --git a/tests/grliq/managers/test_forensic_results.py b/tests/grliq/managers/test_forensic_results.py index a030451..86834d0 100644 --- a/tests/grliq/managers/test_forensic_results.py +++ b/tests/grliq/managers/test_forensic_results.py @@ -1,18 +1,21 @@ from __future__ import annotations +from collections.abc import Callable from typing import TYPE_CHECKING if TYPE_CHECKING: - from generalresearch.grliq.managers.forensic_data import GrlIqDataManager from generalresearch.grliq.managers.forensic_results import ( GrlIqCategoryResultsReader, ) + from generalresearch.grliq.models.forensic_data import GrlIqData class TestGrlIqCategoryResultsReader: def test_filter_category_results( - self, grliq_dm: GrlIqDataManager, grliq_crr: GrlIqCategoryResultsReader + self, + grliq_data_factory: Callable[..., GrlIqData], + grliq_crr: GrlIqCategoryResultsReader, ): from generalresearch.grliq.models.forensic_result import ( GrlIqForensicCategoryResult, @@ -20,8 +23,8 @@ class TestGrlIqCategoryResultsReader: ) # this is just testing that it doesn't fail - grliq_dm.create_dummy(is_attempt_allowed=True) - grliq_dm.create_dummy(is_attempt_allowed=True) + grliq_data_factory(is_attempt_allowed=True) + grliq_data_factory(is_attempt_allowed=True) res = grliq_crr.filter_category_results(limit=2, phase=Phase.OFFERWALL_ENTER)[0] assert res.get("category_result") diff --git a/tests/managers/gr/test_business.py b/tests/managers/gr/test_business.py index 0d5b0d5..6a930b4 100644 --- a/tests/managers/gr/test_business.py +++ b/tests/managers/gr/test_business.py @@ -76,30 +76,30 @@ class TestBusinessManager: assert isinstance(instance, Business) assert isinstance(instance.id, int) - def test_get_or_create(self, business_manager: BusinessManager): + def test_get_or_create(self, gr_business_manager: BusinessManager): uuid_key = uuid4().hex - assert business_manager.get_by_uuid(business_uuid=uuid_key) is None + assert gr_business_manager.get_by_uuid(business_uuid=uuid_key) is None - instance = business_manager.get_or_create( + instance = gr_business_manager.get_or_create( uuid=uuid_key, name=f"name-{uuid4().hex[:6]}", ) - res = business_manager.get_by_uuid(business_uuid=uuid_key) + res = gr_business_manager.get_by_uuid(business_uuid=uuid_key) assert isinstance(res, Business) assert res.id == instance.id def test_get_all( self, - business_manager: BusinessManager, + gr_business_manager: BusinessManager, gr_business_factory: Callable[..., Business], ): - res1 = business_manager.get_all() + res1 = gr_business_manager.get_all() assert isinstance(res1, list) gr_business_factory() - res2 = business_manager.get_all() + res2 = gr_business_manager.get_all() assert len(res1) == len(res2) - 1 @pytest.mark.skip(reason="TODO") @@ -108,42 +108,42 @@ class TestBusinessManager: def test_get_by_user_id( self, - business_manager: BusinessManager, + gr_business_manager: BusinessManager, gr_user: GRUser, team_manager: TeamManager, membership_manager: MembershipManager, gr_business_factory: Callable[..., Business], gr_team_factory: Callable[..., Team], ): - res = business_manager.get_by_user_id(user_id=gr_user.id) + res = gr_business_manager.get_by_user_id(user_id=gr_user.id) assert len(res) == 0 # Create a business: Business, but don't add it to anything b1 = gr_business_factory() - res = business_manager.get_by_user_id(user_id=gr_user.id) + res = gr_business_manager.get_by_user_id(user_id=gr_user.id) assert len(res) == 0 # Create a Team, but don't create any Memberships t1 = gr_team_factory() - res = business_manager.get_by_user_id(user_id=gr_user.id) + res = gr_business_manager.get_by_user_id(user_id=gr_user.id) assert len(res) == 0 # Create a Membership for the gr_user to the Team... but it doesn't # matter because the Team doesn't have any Business yet _ = membership_manager.create(team=t1, gr_user=gr_user) - res = business_manager.get_by_user_id(user_id=gr_user.id) + res = gr_business_manager.get_by_user_id(user_id=gr_user.id) assert len(res) == 0 # Add the Business to the Team... now the Business should be available # to the gr_user team_manager.add_business(team=t1, business=b1) - res = business_manager.get_by_user_id(user_id=gr_user.id) + res = gr_business_manager.get_by_user_id(user_id=gr_user.id) assert len(res) == 1 # Add another Business to the Team! b2 = gr_business_factory() team_manager.add_business(team=t1, business=b2) - res = business_manager.get_by_user_id(user_id=gr_user.id) + res = gr_business_manager.get_by_user_id(user_id=gr_user.id) assert len(res) == 2 @pytest.mark.skip(reason="TODO") @@ -151,14 +151,16 @@ class TestBusinessManager: pass def test_get_by_uuid( - self, gr_business: Business, business_manager: BusinessManager + self, gr_business: Business, gr_business_manager: BusinessManager ): - instance = business_manager.get_by_uuid(business_uuid=gr_business.uuid) + instance = gr_business_manager.get_by_uuid(business_uuid=gr_business.uuid) assert isinstance(instance, Business) assert gr_business.id == instance.id - def test_get_by_id(self, gr_business: Business, business_manager: BusinessManager): - instance = business_manager.get_by_id(business_id=gr_business.id) + def test_get_by_id( + self, gr_business: Business, gr_business_manager: BusinessManager + ): + instance = gr_business_manager.get_by_id(business_id=gr_business.id) assert isinstance(instance, Business) assert gr_business.uuid == instance.uuid diff --git a/tests/managers/thl/test_contest/test_milestone.py b/tests/managers/thl/test_contest/test_milestone.py index dbb2016..dab02e7 100644 --- a/tests/managers/thl/test_contest/test_milestone.py +++ b/tests/managers/thl/test_contest/test_milestone.py @@ -6,10 +6,10 @@ from typing import TYPE_CHECKING from generalresearch.models.thl.contest.definitions import ( ContestEndReason, + ContestEntryTrigger, ContestStatus, ) from generalresearch.models.thl.contest.milestone import ( - ContestEntryTrigger, MilestoneContest, MilestoneUserView, ) diff --git a/tests/managers/thl/test_contest/test_raffle.py b/tests/managers/thl/test_contest/test_raffle.py index 7803952..0b2b852 100644 --- a/tests/managers/thl/test_contest/test_raffle.py +++ b/tests/managers/thl/test_contest/test_raffle.py @@ -17,17 +17,17 @@ from generalresearch.models.thl.contest import ( ContestEntryRule, ContestPrize, ) +from generalresearch.models.thl.contest.contest_entry import ( + ContestEntry, + ContestEntryType, +) from generalresearch.models.thl.contest.definitions import ( ContestEndReason, ContestPrizeKind, ContestStatus, ) from generalresearch.models.thl.contest.exceptions import ContestError -from generalresearch.models.thl.contest.raffle import ( - ContestEntry, - ContestEntryType, - RaffleContest, -) +from generalresearch.models.thl.contest.raffle import RaffleContest if TYPE_CHECKING: from generalresearch.managers.thl.contest_manager import ContestManager diff --git a/tests/managers/thl/test_ipinfo.py b/tests/managers/thl/test_ipinfo.py index 6954163..47b1712 100644 --- a/tests/managers/thl/test_ipinfo.py +++ b/tests/managers/thl/test_ipinfo.py @@ -31,14 +31,16 @@ class TestIPGeonameManager: assert isinstance(instance, IPGeonameManager) assert isinstance(ip_geoname_manager, IPGeonameManager) - def test_create(self, ip_geoname_manager: IPGeonameManager): - - instance = ip_geoname_manager.create_dummy() + def test_create( + self, + ip_geoname_factory: Callable[..., IPGeoname], + ip_geoname_manager: IPGeonameManager, + ): + instance = ip_geoname_factory() assert isinstance(instance, IPGeoname) res = ip_geoname_manager.fetch_geoname_ids(filter_ids=[instance.geoname_id]) - assert res[0].model_dump_json() == instance.model_dump_json() @@ -51,13 +53,15 @@ class TestIPInformationManager: assert isinstance(instance, IPInformationManager) assert isinstance(ip_information_manager, IPInformationManager) - def test_create(self, ip_information_manager: IPInformationManager): - instance = ip_information_manager.create_dummy() - + def test_create( + self, + ip_geoname_factory: Callable[..., IPGeoname], + ip_information_manager: IPInformationManager, + ): + instance = ip_geoname_factory() assert isinstance(instance, IPInformation) res = ip_information_manager.fetch_ip_information(filter_ips=[instance.ip]) - assert res[0].model_dump_json() == instance.model_dump_json() def test_prefetch_geoname( diff --git a/tests/managers/thl/test_ledger/test_lm_accounts.py b/tests/managers/thl/test_ledger/test_lm_accounts.py index 57a2261..3af10e7 100644 --- a/tests/managers/thl/test_ledger/test_lm_accounts.py +++ b/tests/managers/thl/test_ledger/test_lm_accounts.py @@ -14,8 +14,10 @@ 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 UUIDStr from generalresearch.models.thl.ledger import ( + AccountType, + Direction, LedgerAccount, LedgerEntry, ) 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 9ecc1bc..166598e 100644 --- a/tests/managers/thl/test_ledger/test_lm_tx_locks.py +++ b/tests/managers/thl/test_ledger/test_lm_tx_locks.py @@ -4,10 +4,10 @@ import logging 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 pytest import LogCaptureFixture from generalresearch.managers.thl.ledger_manager.conditions import ( generate_condition_mp_payment, 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 b0484ae..cda88da 100644 --- a/tests/managers/thl/test_ledger/test_thl_lm_tx.py +++ b/tests/managers/thl/test_ledger/test_thl_lm_tx.py @@ -128,10 +128,11 @@ class TestThlLedgerTxManager: thl_ledger_manager: ThlLedgerManager, ledger_manager: LedgerManager, session_manager: SessionManager, + product_factory: Callable[..., Product], ): delete_ledger_db() create_main_accounts() - product = product_manager.create_dummy( + product = product_factory( payout_config=PayoutConfig( payout_transformation=PayoutTransformation( f="payout_transformation_amt" diff --git a/tests/managers/thl/test_ledger/test_wallet.py b/tests/managers/thl/test_ledger/test_wallet.py index 1ee9bf9..dc1feec 100644 --- a/tests/managers/thl/test_ledger/test_wallet.py +++ b/tests/managers/thl/test_ledger/test_wallet.py @@ -22,8 +22,10 @@ if TYPE_CHECKING: @pytest.fixture() -def schrute_product(product_manager: ProductManager) -> Product: - return product_manager.create_dummy( +def schrute_product( + product_factory: Callable[..., Product], product_manager: ProductManager +) -> Product: + return product_factory( user_wallet_config=UserWalletConfig(enabled=True, amt=False), payout_config=PayoutConfig( payout_transformation=PayoutTransformation( diff --git a/tests/managers/thl/test_product.py b/tests/managers/thl/test_product.py index 644dc90..81e0122 100644 --- a/tests/managers/thl/test_product.py +++ b/tests/managers/thl/test_product.py @@ -24,8 +24,12 @@ if TYPE_CHECKING: class TestProductManagerGetMethods: - def test_get_by_uuid(self, product_manager: ProductManager): - product: Product = product_manager.create_dummy( + def test_get_by_uuid( + self, + product_manager: ProductManager, + product_factory: Callable[..., Product], + ): + product: Product = product_factory( product_id=uuid4().hex, team_id=uuid4().hex, name=f"Test Product ID #{uuid4().hex[:6]}", @@ -44,12 +48,14 @@ class TestProductManagerGetMethods: product_manager.get_by_uuid(product_uuid=uuid4().hex) assert "product not found" in str(cm.value) - def test_get_by_uuids(self, product_manager: ProductManager): + def test_get_by_uuids( + self, product_factory: Callable[..., Product], product_manager: ProductManager + ): cnt = 5 product_uuids = [uuid4().hex for _ in range(cnt)] for product_id in product_uuids: - product_manager.create_dummy( + product_factory( product_id=product_id, team_id=uuid4().hex, name=f"Test Product ID #{uuid4().hex[:6]}", @@ -69,8 +75,10 @@ class TestProductManagerGetMethods: product_manager.get_by_uuids(product_uuids=product_uuids + ["abc123"]) assert "invalid uuid" in str(cm.value) - def test_get_by_uuid_if_exists(self, product_manager: ProductManager): - product: Product = product_manager.create_dummy( + def test_get_by_uuid_if_exists( + self, product_factory: Callable[..., Product], product_manager: ProductManager + ): + product: Product = product_factory( product_id=uuid4().hex, team_id=uuid4().hex, name=f"Test Product ID #{uuid4().hex[:6]}", @@ -81,10 +89,12 @@ class TestProductManagerGetMethods: instance = product_manager.get_by_uuid_if_exists(product_uuid="abc123") assert instance == None - def test_get_by_uuids_if_exists(self, product_manager: ProductManager): + def test_get_by_uuids_if_exists( + self, product_factory: Callable[..., Product], product_manager: ProductManager + ): product_uuids = [uuid4().hex for _ in range(2)] for product_id in product_uuids: - product_manager.create_dummy( + product_factory( product_id=product_id, team_id=uuid4().hex, name=f"Test Product ID #{uuid4().hex[:6]}", @@ -113,13 +123,15 @@ class TestProductManagerGetMethods: # for instance in res: # assert isinstance(instance, Product) - def test_get_by_business_ids(self, product_manager: ProductManager): + def test_get_by_business_ids( + self, product_factory: Callable[..., Product], product_manager: ProductManager + ): business_ids = [uuid4().hex for _ in range(5)] product_manager.fetch_uuids(business_uuids=business_ids) for business_id in business_ids: - product_manager.create( + product_factory( product_id=uuid4().hex, team_id=None, business_id=business_id, @@ -131,8 +143,10 @@ class TestProductManagerGetMethods: class TestProductManagerCreation: - def test_base(self, product_manager: ProductManager): - instance = product_manager.create_dummy( + def test_base( + self, product_factory: Callable[..., Product], product_manager: ProductManager + ): + instance = product_factory( product_id=uuid4().hex, team_id=uuid4().hex, name=f"New Test Product {uuid4().hex[:6]}", @@ -235,10 +249,12 @@ class TestProductManager: assert instance.user_create_config.max_hourly_create_limit is None assert not instance.user_wallet_config.enabled - def test_sources(self, product_manager: ProductManager): + def test_sources( + self, product_factory: Callable[..., Product], product_manager: ProductManager + ): user_defined = [SourceConfig(name=Source.DYNATA, active=False)] sources_config = SourcesConfig(user_defined=user_defined) - p = product_manager.create_dummy(sources_config=sources_config) + p = product_factory(sources_config=sources_config) p2 = product_manager.get_by_uuid(p.id) @@ -250,7 +266,9 @@ class TestProductManager: assert not dynata.active assert all(x.active is True for x in p2.sources if x.name != Source.DYNATA) - def test_global_sources(self, product_manager: ProductManager): + def test_global_sources( + self, product_factory: Callable[..., Product], product_manager: ProductManager + ): sources_config = SupplyConfig( policies=[ SupplyPolicy( @@ -261,7 +279,7 @@ class TestProductManager: ) ] ) - p1 = product_manager.create_dummy(sources_config=sources_config) + p1 = product_factory(sources_config=sources_config) p2 = product_manager.get_by_uuid(p1.id) assert p1 == p2 @@ -277,8 +295,10 @@ class TestProductManager: p2 = product_manager.get_by_uuid(p1.id) assert p1 == p2 - def test_user_health_config(self, product_manager: ProductManager): - p = product_manager.create_dummy( + def test_user_health_config( + self, product_factory: Callable[..., Product], product_manager: ProductManager + ): + p = product_factory( user_health_config=UserHealthConfig(banned_countries=["ng", "in"]) ) @@ -288,10 +308,10 @@ class TestProductManager: assert p2.user_health_config.banned_countries == ["in", "ng"] assert p2.user_health_config.allow_ban_iphist - def test_profiling_config(self, product_manager: ProductManager): - p = product_manager.create_dummy( - profiling_config=ProfilingConfig(max_questions=1) - ) + def test_profiling_config( + self, product_factory: Callable[..., Product], product_manager: ProductManager + ): + p = product_factory(profiling_config=ProfilingConfig(max_questions=1)) p2 = product_manager.get_by_uuid(p.id) assert p == p2 @@ -335,8 +355,10 @@ class TestProductManager: class TestProductManagerUpdate: - def test_update(self, product_manager: ProductManager): - p = product_manager.create_dummy() + def test_update( + self, product_factory: Callable[..., Product], product_manager: ProductManager + ): + p = product_factory() p.name = "new name" p.enabled = False p.user_create_config = UserCreateConfig(min_hourly_create_limit=200) @@ -356,8 +378,10 @@ class TestProductManagerUpdate: class TestProductManagerCacheClear: - def test_cache_clear(self, product_manager: ProductManager): - p = product_manager.create_dummy() + def test_cache_clear( + self, product_factory: Callable[..., Product], product_manager: ProductManager + ): + p = product_factory() product_manager.get_by_uuid(product_uuid=p.id) product_manager.get_by_uuid(product_uuid=p.id) product_manager.pg_config.execute_write( diff --git a/tests/managers/thl/test_task_status.py b/tests/managers/thl/test_task_status.py index 9846ce0..4a401fa 100644 --- a/tests/managers/thl/test_task_status.py +++ b/tests/managers/thl/test_task_status.py @@ -40,18 +40,22 @@ finish3 = start3 + timedelta(minutes=5) @pytest.fixture(scope="session") -def bp1(product_manager: ProductManager) -> Product: +def bp1( + product_factory: Callable[..., Product], product_manager: ProductManager +) -> Product: # user wallet disabled, payout xform NULL - return product_manager.create_dummy( + return product_factory( user_wallet_config=UserWalletConfig(enabled=False), payout_config=PayoutConfig(), ) @pytest.fixture(scope="session") -def bp2(product_manager: ProductManager) -> Product: +def bp2( + product_factory: Callable[..., Product], product_manager: ProductManager +) -> Product: # user wallet disabled, payout xform 40% - return product_manager.create_dummy( + return product_factory( user_wallet_config=UserWalletConfig(enabled=False), payout_config=PayoutConfig( payout_transformation=PayoutTransformation( @@ -63,9 +67,11 @@ def bp2(product_manager: ProductManager) -> Product: @pytest.fixture(scope="session") -def bp3(product_manager: ProductManager) -> Product: +def bp3( + product_factory: Callable[..., Product], product_manager: ProductManager +) -> Product: # user wallet enabled, payout xform 50% - return product_manager.create_dummy( + return product_factory( user_wallet_config=UserWalletConfig(enabled=True), payout_config=PayoutConfig( payout_transformation=PayoutTransformation( diff --git a/tests/managers/thl/test_user_manager/test_base.py b/tests/managers/thl/test_user_manager/test_base.py index 4a9750e..c69f297 100644 --- a/tests/managers/thl/test_user_manager/test_base.py +++ b/tests/managers/thl/test_user_manager/test_base.py @@ -1,4 +1,5 @@ import logging +from collections.abc import Callable from datetime import UTC, datetime from random import randint from typing import TYPE_CHECKING @@ -7,9 +8,11 @@ from uuid import uuid4 import pytest from generalresearch.managers.thl.user_manager import ( - UserCreateNotAllowedError, get_bp_user_create_limit_hourly, ) +from generalresearch.managers.thl.user_manager.exceptions import ( + UserCreateNotAllowedError, +) from generalresearch.managers.thl.user_manager.mysql_user_manager import ( MysqlUserManager, ) @@ -152,11 +155,11 @@ class TestCreateUserManager: def test_create_user( self, - product_manager: ProductManager, + product_factory: Callable[..., Product], thl_web_rw: PostgresConfig, user_manager: UserManager, ): - product: Product = product_manager.create_dummy( + product: Product = product_factory( user_create_config=UserCreateConfig( min_hourly_create_limit=10, max_hourly_create_limit=69 ), @@ -195,11 +198,11 @@ class TestCreateUserManager: def test_create_user_integrity_error( self, - product_manager: ProductManager, user_manager: UserManager, + product_factory: Callable[..., Product], caplog, ): - product: Product = product_manager.create_dummy( + product: Product = product_factory( product_id=uuid4().hex, team_id=uuid4().hex, name=f"Test Product ID #{uuid4().hex[:6]}", @@ -241,10 +244,13 @@ class TestCreateUserManager: assert user1 == user2 def test_raise_allow_user_create( - self, product_manager: ProductManager, user_manager: UserManager + self, + product_manager: ProductManager, + user_manager: UserManager, + product_factory: Callable[..., Product], ): rand_num = randint(25, 200) - product: Product = product_manager.create_dummy( + product: Product = product_factory( product_id=uuid4().hex, team_id=uuid4().hex, name=f"Test Product ID #{uuid4().hex[:6]}", diff --git a/tests/managers/thl/test_wall_manager.py b/tests/managers/thl/test_wall_manager.py index 3215de8..58de7a2 100644 --- a/tests/managers/thl/test_wall_manager.py +++ b/tests/managers/thl/test_wall_manager.py @@ -10,7 +10,7 @@ import pytest from pydantic import PositiveInt from generalresearch.models.definitions import Source -from generalresearch.models.thl.session import ( +from generalresearch.models.thl.definitions import ( ReportValue, Status, StatusCode1, diff --git a/tests/models/gr/test_business.py b/tests/models/gr/test_business.py index 4e0b4e1..e942be5 100644 --- a/tests/models/gr/test_business.py +++ b/tests/models/gr/test_business.py @@ -551,7 +551,6 @@ class TestBusinessBalance: ledger_manager: LedgerManager, product_manager: ProductManager, start: datetime, - thl_web_rr: PostgresConfig, session_with_tx_factory: Callable[..., Session], delete_ledger_db: Callable[..., None], create_main_accounts: Callable[..., None], @@ -964,7 +963,6 @@ class TestBusinessBalance: ledger_manager: LedgerManager, product_manager: ProductManager, start: datetime, - thl_web_rr: PostgresConfig, payout_event_manager, session_with_tx_factory: Callable[..., None], delete_ledger_db: Callable[..., None], @@ -1194,7 +1192,6 @@ class TestBusinessMethods: def test_set_cache( self, gr_business: Business, - gr_db: PostgresConfig, thl_web_rr: PostgresConfig, client_no_amm: DaskClient, mnt_filepath: GRLDatasets, diff --git a/tests/models/gr/test_team.py b/tests/models/gr/test_team.py index a94d53f..0ca9b11 100644 --- a/tests/models/gr/test_team.py +++ b/tests/models/gr/test_team.py @@ -113,13 +113,13 @@ class TestTeam: assert gr_team.businesses is None - gr_team.prefetch_businesses(business_manager=gr_business_manager) + gr_team.prefetch_businesses(gr_business_manager=gr_business_manager) assert isinstance(gr_team.businesses, list) assert len(gr_team.businesses) == 0 team_manager.add_business(team=gr_team, business=business) assert len(gr_team.businesses) == 0 - gr_team.prefetch_businesses(business_manager=gr_business_manager) + gr_team.prefetch_businesses(gr_business_manager=gr_business_manager) assert len(gr_team.businesses) == 1 assert isinstance(gr_team.businesses[0], Business) assert gr_team.businesses[0].uuid == business.uuid @@ -163,12 +163,19 @@ class TestTeamMethods: mnt_gr_api_dir: Path, enriched_wall_merge: EnrichedWallMerge, enriched_session_merge: EnrichedSessionMerge, + product_manager: ProductManager, + gr_user_manager: GRUserManager, + gr_business_manager: BusinessManager, + gr_membership_manager: MembershipManager, ): client = gr_redis_config.create_redis_client() assert client.get(name=gr_team.cache_key) is None gr_team.set_cache( - pg_config=gr_db, + product_manager=product_manager, + gr_user_manager=gr_user_manager, + gr_business_manager=gr_business_manager, + gr_membership_manager=gr_membership_manager, thl_web_rr=thl_web_rr, redis_config=gr_redis_config, client=client_no_amm, @@ -193,6 +200,10 @@ class TestTeamMethods: mnt_gr_api_dir: Path, enriched_wall_merge: EnrichedWallMerge, enriched_session_merge: EnrichedSessionMerge, + product_manager: ProductManager, + gr_user_manager: GRUserManager, + gr_business_manager: BusinessManager, + gr_membership_manager: MembershipManager, ): from generalresearch.models.gr.team import Team @@ -200,7 +211,10 @@ class TestTeamMethods: membership_factory(team=gr_team, gr_user=gr_user) gr_team.set_cache( - pg_config=gr_db, + product_manager=product_manager, + gr_user_manager=gr_user_manager, + gr_business_manager=gr_business_manager, + gr_membership_manager=gr_membership_manager, thl_web_rr=thl_web_rr, redis_config=gr_redis_config, client=client_no_amm, @@ -239,6 +253,7 @@ class TestTeamMethods: mnt_filepath: GRLDatasets, mnt_gr_api_dir: Path, gr_team: Team, + product_manager: ProductManager, ): delete_df_collection(coll=wall_collection) @@ -267,7 +282,7 @@ class TestTeamMethods: ) gr_team.prebuild_enriched_session_parquet( - thl_pg_config=thl_web_rr, + product_manager=product_manager, ds=mnt_filepath, client=client_no_amm, mnt_gr_api=mnt_gr_api_dir, @@ -295,6 +310,7 @@ class TestTeamMethods: mnt_filepath: GRLDatasets, mnt_gr_api_dir: Path, gr_team: Team, + product_manager: ProductManager, ): delete_df_collection(coll=wall_collection) @@ -323,7 +339,7 @@ class TestTeamMethods: ) gr_team.prebuild_enriched_wall_parquet( - thl_pg_config=thl_web_rr, + product_manager=product_manager, ds=mnt_filepath, client=client_no_amm, mnt_gr_api=mnt_gr_api_dir, diff --git a/tests/models/thl/test_product.py b/tests/models/thl/test_product.py index 446b59f..f1050bb 100644 --- a/tests/models/thl/test_product.py +++ b/tests/models/thl/test_product.py @@ -64,9 +64,11 @@ class TestProduct: # We're not excluding anything here, only in the "*Out" variants assert "id_int" in res - def test_init_db(self, product_manager: ProductManager): + def test_init_db( + self, product_factory: Callable[..., Product], product_manager: ProductManager + ): # By default, just a Pydantic instance doesn't have an id_int - instance = product_manager.create_dummy() + instance = product_factory() assert isinstance(instance.id_int, int) res = instance.model_dump_json() diff --git a/tests/models/thl/test_user.py b/tests/models/thl/test_user.py index bc941d4..68b413c 100644 --- a/tests/models/thl/test_user.py +++ b/tests/models/thl/test_user.py @@ -18,6 +18,7 @@ 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 + from generalresearch.models.thl.userhealth import AuditLog class TestUserUserID: @@ -621,12 +622,17 @@ class TestUserSerialization: class TestUserMethods: - def test_audit_log(self, user: User, audit_log_manager: AuditLogManager): + def test_audit_log( + self, + audit_log_factory: Callable[..., AuditLog], + user: User, + audit_log_manager: AuditLogManager, + ): assert user.audit_log is None user.prefetch_audit_log(audit_log_manager=audit_log_manager) assert user.audit_log == [] - audit_log_manager.create_dummy(user_id=user.user_id) + audit_log_factory(user_id=user.user_id) user.prefetch_audit_log(audit_log_manager=audit_log_manager) assert len(user.audit_log) == 1 -- 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') 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 From 3c4fdaf7804999bcadc32bc6bb6fce2ad0435a61 Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Thu, 3 Sep 2026 10:37:19 -0700 Subject: fixture names --- test_utils/models/thl/conftest.py | 1 - tests/managers/gr/test_business.py | 4 ++-- tests/managers/gr/test_team.py | 6 ++++-- tests/managers/thl/test_wall_manager.py | 2 -- tests/models/gr/test_team.py | 6 +++--- 5 files changed, 9 insertions(+), 10 deletions(-) (limited to 'test_utils/models') diff --git a/test_utils/models/thl/conftest.py b/test_utils/models/thl/conftest.py index 5dc46cd..021b19e 100644 --- a/test_utils/models/thl/conftest.py +++ b/test_utils/models/thl/conftest.py @@ -42,7 +42,6 @@ if TYPE_CHECKING: 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 diff --git a/tests/managers/gr/test_business.py b/tests/managers/gr/test_business.py index 6a930b4..70f50ca 100644 --- a/tests/managers/gr/test_business.py +++ b/tests/managers/gr/test_business.py @@ -111,7 +111,7 @@ class TestBusinessManager: gr_business_manager: BusinessManager, gr_user: GRUser, team_manager: TeamManager, - membership_manager: MembershipManager, + gr_membership_manager: MembershipManager, gr_business_factory: Callable[..., Business], gr_team_factory: Callable[..., Team], ): @@ -130,7 +130,7 @@ class TestBusinessManager: # Create a Membership for the gr_user to the Team... but it doesn't # matter because the Team doesn't have any Business yet - _ = membership_manager.create(team=t1, gr_user=gr_user) + _ = gr_membership_manager.create(team=t1, gr_user=gr_user) res = gr_business_manager.get_by_user_id(user_id=gr_user.id) assert len(res) == 0 diff --git a/tests/managers/gr/test_team.py b/tests/managers/gr/test_team.py index 751e33c..878a9ca 100644 --- a/tests/managers/gr/test_team.py +++ b/tests/managers/gr/test_team.py @@ -18,8 +18,10 @@ if TYPE_CHECKING: class TestMembershipManager: - def test_init(self, membership_manager: MembershipManager, gr_db: PostgresConfig): - assert membership_manager.pg_config == gr_db + def test_init( + self, gr_membership_manager: MembershipManager, gr_db: PostgresConfig + ): + assert gr_membership_manager.pg_config == gr_db class TestTeamManager: diff --git a/tests/managers/thl/test_wall_manager.py b/tests/managers/thl/test_wall_manager.py index 70db71e..29d2660 100644 --- a/tests/managers/thl/test_wall_manager.py +++ b/tests/managers/thl/test_wall_manager.py @@ -247,8 +247,6 @@ class TestWallCacheManager: def test_get_wall_events( self, wall_cache_manager: WallCacheManager, - wall_manager: WallManager, - session_manager: SessionManager, user: User, session_factory: Callable[..., Session], wall_factory: Callable[..., Wall], diff --git a/tests/models/gr/test_team.py b/tests/models/gr/test_team.py index 0ca9b11..8ebedb6 100644 --- a/tests/models/gr/test_team.py +++ b/tests/models/gr/test_team.py @@ -82,7 +82,7 @@ class TestTeam: self, gr_team: Team, gr_user_factory: Callable[..., GRUser], - membership_manager: MembershipManager, + gr_membership_manager: MembershipManager, gr_user_manager: GRUserManager, ): assert gr_team.gr_users is None @@ -92,13 +92,13 @@ class TestTeam: assert len(gr_team.gr_users) == 0 # Create a new Membership - membership_manager.create(team=gr_team, gr_user=gr_user_factory()) + gr_membership_manager.create(team=gr_team, gr_user=gr_user_factory()) assert len(gr_team.gr_users) == 0 gr_team.prefetch_gr_users(gr_user_manager=gr_user_manager) assert len(gr_team.gr_users) == 1 # Create another Membership - membership_manager.create(team=gr_team, gr_user=gr_user_factory()) + gr_membership_manager.create(team=gr_team, gr_user=gr_user_factory()) assert len(gr_team.gr_users) == 1 gr_team.prefetch_gr_users(gr_user_manager=gr_user_manager) assert len(gr_team.gr_users) == 2 -- cgit v1.2.3 From 48863e9431d50fd86405036b55b333b7734b00ac Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Thu, 3 Sep 2026 11:30:55 -0700 Subject: import fixes for <3.13 --- test_utils/models/network/conftest.py | 4 +--- test_utils/models/thl/conftest.py | 2 +- test_utils/models/upk/conftest.py | 6 ++---- 3 files changed, 4 insertions(+), 8 deletions(-) (limited to 'test_utils/models') diff --git a/test_utils/models/network/conftest.py b/test_utils/models/network/conftest.py index 4ff59ee..c7bcc7e 100644 --- a/test_utils/models/network/conftest.py +++ b/test_utils/models/network/conftest.py @@ -24,9 +24,7 @@ from generalresearch.models.network.tool_run_command import ( RDNSRunCommand, RDNSRunCommandOptions, ) - -if TYPE_CHECKING: - from generalresearch.pg_helper import PostgresConfig +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 021b19e..f1cb785 100644 --- a/test_utils/models/thl/conftest.py +++ b/test_utils/models/thl/conftest.py @@ -32,6 +32,7 @@ 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 +from generalresearch.pg_helper import PostgresConfig if TYPE_CHECKING: from generalresearch.currency import USDCent @@ -73,7 +74,6 @@ if TYPE_CHECKING: from generalresearch.models.thl.user_iphistory import IPRecord from generalresearch.models.thl.userhealth import AuditLog from generalresearch.models.thl.wallet.cashout_method import CashMailOrderData - from generalresearch.pg_helper import PostgresConfig fake = faker.Faker() diff --git a/test_utils/models/upk/conftest.py b/test_utils/models/upk/conftest.py index 59266b2..ad96bbb 100644 --- a/test_utils/models/upk/conftest.py +++ b/test_utils/models/upk/conftest.py @@ -3,15 +3,13 @@ from __future__ import annotations import os import time from collections.abc import Callable -from typing import TYPE_CHECKING from uuid import UUID import pandas as pd import pytest -if TYPE_CHECKING: - from generalresearch.managers.thl.category import CategoryManager - from generalresearch.pg_helper import PostgresConfig +from generalresearch.managers.thl.category import CategoryManager +from generalresearch.pg_helper import PostgresConfig def insert_data_from_csv( -- cgit v1.2.3 From 4415e04365e4a92b5b7d7ee55fa5193692d1d29b Mon Sep 17 00:00:00 2001 From: stuppie Date: Thu, 3 Sep 2026 16:22:10 -0600 Subject: theres 2 fixtures called session_factory that do different things. this took me a while to figure out --- test_utils/models/conftest.py | 165 -------------------------------------- test_utils/models/thl/conftest.py | 139 ++++++++++++++++++++++++++++++-- 2 files changed, 134 insertions(+), 170 deletions(-) (limited to 'test_utils/models') diff --git a/test_utils/models/conftest.py b/test_utils/models/conftest.py index 9edadd3..43f18c1 100644 --- a/test_utils/models/conftest.py +++ b/test_utils/models/conftest.py @@ -1,192 +1,27 @@ from __future__ import annotations from collections.abc import Callable -from datetime import UTC, datetime, timedelta -from decimal import Decimal -from random import choice as randchoice -from random import randint from typing import TYPE_CHECKING from uuid import uuid4 import pytest -from pydantic import AwareDatetime, PositiveInt -from pytest import FixtureRequest from pytest import FixtureRequest as Request 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 if TYPE_CHECKING: - from generalresearch.currency import USDCent from generalresearch.managers.thl.buyer import BuyerManager - from generalresearch.managers.thl.ipinfo import ( - IPGeonameManager, - IPInformationManager, - ) - from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager - from generalresearch.managers.thl.payout import ( - BusinessPayoutEventManager, - ) from generalresearch.managers.thl.product import ProductManager - from generalresearch.managers.thl.session import SessionManager from generalresearch.managers.thl.survey import SurveyManager - from generalresearch.managers.thl.user_manager.user_manager import UserManager - from generalresearch.managers.thl.userhealth import AuditLogManager, IPRecordManager - from generalresearch.managers.thl.wall import WallManager - from generalresearch.models.gr.business import ( - Business, - ) - from generalresearch.models.gr.team import Team - from generalresearch.models.thl.ipinfo import IPGeoname, IPInformation - from generalresearch.models.thl.payout import ( - BrokerageProductPayoutEvent, - ) from generalresearch.models.thl.product import ( PayoutConfig, Product, ) - from generalresearch.models.thl.session import Session, Wall - from generalresearch.models.thl.user import User # === THL === -@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 - - def _inner( - user: User, - # Wall details - wall_count: int = 5, - wall_req_cpi: Decimal = Decimal(".50"), - wall_req_cpis: list[Decimal] | None = None, - wall_statuses: list[Status] | None = None, - wall_source: Source = Source.TESTING, - # Session details - final_status: Status = Status.COMPLETE, - started: datetime = utc_hour_ago, - ) -> Session: - if wall_req_cpis: - assert len(wall_req_cpis) == wall_count - if wall_statuses: - assert len(wall_statuses) == wall_count - - 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 - wall_started = s.started + timedelta(milliseconds=1) - else: - # Subsequent Wall events - last_wall = s.wall_events[-1] - assert last_wall.finished, "Can't add new Walls until prior finishes" - wall_started = last_wall.started + timedelta(milliseconds=1) - - w = wall_factory( - session_id=s.id, - source=wall_source, - user_id=s.user_id, - started=wall_started, - req_cpi=wall_req_cpis[idx] if wall_req_cpis else wall_req_cpi, - ) - s.append_wall_event(w=w) - - # If it's the last wall in the session, respect the final_status - # value for the Session - if wall_statuses: - _final_status = wall_statuses[idx] - else: - _final_status = final_status if idx == wall_count - 1 else Status.FAIL - - options = list(WALL_ALLOWED_STATUS_STATUS_CODE.get(_final_status, {})) - wall_manager.finish( - wall=w, - status=_final_status, - status_code_1=randchoice(options), - finished=w.started + timedelta(seconds=randint(a=60 * 2, b=60 * 10)), - ) - - return s - - return _inner - - -@pytest.fixture(scope="function") -def finished_session_factory( - session_factory: Callable[..., Session], - session_manager: SessionManager, - utc_hour_ago: datetime, -) -> Callable[..., Session]: - from generalresearch.models.thl.session import Source - - def _inner( - user: User, - # Wall details - wall_count: int = 5, - wall_req_cpi: Decimal = Decimal(".50"), - wall_req_cpis: list[Decimal] | None = None, - wall_statuses: list[Status] | None = None, - wall_source: Source = Source.TESTING, - # Session details - final_status: Status = Status.COMPLETE, - started: datetime = utc_hour_ago, - ) -> Session: - s: Session = session_factory( - user=user, - wall_count=wall_count, - wall_req_cpi=wall_req_cpi, - wall_req_cpis=wall_req_cpis, - wall_statuses=wall_statuses, - wall_source=wall_source, - final_status=final_status, - started=started, - ) - status, status_code_1 = s.determine_session_status() - _, _, bp_pay, user_pay = s.determine_payments() - session_manager.finish_with_status( - s, - finished=s.wall_events[-1].finished, - payout=bp_pay, - user_payout=user_pay, - status=status, - status_code_1=status_code_1, - ) - return s - - return _inner - - -@pytest.fixture -def session( - user: User, - session_manager: SessionManager, - wall_manager: WallManager, - session_factory: Callable[..., Session], - wall_factory: Callable[..., Wall], -) -> Session: - - 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, - ) - session.append_wall_event(w=wall) - - return session - - @pytest.fixture() def payout_config(request: Request) -> PayoutConfig: from generalresearch.models.thl.product import ( diff --git a/test_utils/models/thl/conftest.py b/test_utils/models/thl/conftest.py index f1cb785..6f2835e 100644 --- a/test_utils/models/thl/conftest.py +++ b/test_utils/models/thl/conftest.py @@ -4,6 +4,7 @@ from collections.abc import Callable from datetime import UTC, datetime, timedelta from decimal import ROUND_DOWN, Decimal from random import choice as rand_choice +from random import choice as randchoice from random import randint, random from typing import TYPE_CHECKING, Any from uuid import uuid4 @@ -98,7 +99,7 @@ def wall_factory( ) -> Callable[..., Wall]: def _inner( - wall_status: Status, + wall_status: Status = Status.FAIL, save: bool = True, session: Session | None = None, session_id: PositiveInt | None = None, @@ -113,7 +114,6 @@ def wall_factory( """To be used in tests, where we don't care about certain fields""" if save: - 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), @@ -128,9 +128,9 @@ def wall_factory( 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" + 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 @@ -255,11 +255,139 @@ def session(session_factory: Callable[..., Session]) -> Session: return session_factory(save=True) +@pytest.fixture +def session_w_wall( + user: User, + session: Session, + wall_factory: Callable[..., Wall], +) -> Session: + + wall: Wall = wall_factory( + session_id=session.id, + user_id=session.user_id, + started=session.started, + ) + session.append_wall_event(w=wall) + + return session + + @pytest.fixture() def unsaved_session(session_factory: Callable[..., Session]) -> Session: return session_factory(save=False) +@pytest.fixture +def session_w_wall_factory( + wall_manager: WallManager, + utc_hour_ago: datetime, + session_factory: Callable[..., Session], + wall_factory: Callable[..., Wall], +) -> Callable[..., Session]: + from generalresearch.models.thl.session import Source + + def _inner( + user: User, + # Wall details + wall_count: int = 5, + wall_req_cpi: Decimal = Decimal(".50"), + wall_req_cpis: list[Decimal] | None = None, + wall_statuses: list[Status] | None = None, + wall_source: Source = Source.TESTING, + # Session details + final_status: Status = Status.COMPLETE, + started: datetime = utc_hour_ago, + ) -> Session: + if wall_req_cpis: + assert len(wall_req_cpis) == wall_count + if wall_statuses: + assert len(wall_statuses) == wall_count + + 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 + wall_started = s.started + timedelta(milliseconds=1) + else: + # Subsequent Wall events + last_wall = s.wall_events[-1] + assert last_wall.finished, "Can't add new Walls until prior finishes" + wall_started = last_wall.started + timedelta(milliseconds=1) + + w = wall_factory( + session_id=s.id, + source=wall_source, + user_id=s.user_id, + started=wall_started, + req_cpi=wall_req_cpis[idx] if wall_req_cpis else wall_req_cpi, + ) + s.append_wall_event(w=w) + + # If it's the last wall in the session, respect the final_status + # value for the Session + if wall_statuses: + _final_status = wall_statuses[idx] + else: + _final_status = final_status if idx == wall_count - 1 else Status.FAIL + + options = list(WALL_ALLOWED_STATUS_STATUS_CODE.get(_final_status, {})) + wall_manager.finish( + wall=w, + status=_final_status, + status_code_1=randchoice(options), + finished=w.started + timedelta(seconds=randint(a=60 * 2, b=60 * 10)), + ) + + return s + + return _inner + + +@pytest.fixture(scope="function") +def finished_session_factory( + session_w_wall_factory: Callable[..., Session], + session_manager: SessionManager, + utc_hour_ago: datetime, +) -> Callable[..., Session]: + from generalresearch.models.thl.session import Source + + def _inner( + user: User, + # Wall details + wall_count: int = 5, + wall_req_cpi: Decimal = Decimal(".50"), + wall_req_cpis: list[Decimal] | None = None, + wall_statuses: list[Status] | None = None, + wall_source: Source = Source.TESTING, + # Session details + final_status: Status = Status.COMPLETE, + started: datetime = utc_hour_ago, + ) -> Session: + s: Session = session_w_wall_factory( + user=user, + wall_count=wall_count, + wall_req_cpi=wall_req_cpi, + wall_req_cpis=wall_req_cpis, + wall_statuses=wall_statuses, + wall_source=wall_source, + final_status=final_status, + started=started, + ) + status, status_code_1 = s.determine_session_status() + _, _, bp_pay, user_pay = s.determine_payments() + session_manager.finish_with_status( + s, + finished=s.wall_events[-1].finished, + payout=bp_pay, + user_payout=user_pay, + status=status, + status_code_1=status_code_1, + ) + return s + + return _inner + + # --- Product --- @@ -524,6 +652,7 @@ def unsaved_ip_record(ip_record_factory: Callable[..., IPRecord]) -> IPRecord: def user_factory( user_manager: UserManager, thl_web_rr: PostgresConfig, + product_factory: Callable[..., Product], ) -> Callable[..., User]: def _inner( -- cgit v1.2.3 From 780bf9ecbe3e444c3d7d6944c5f610f4ef390ca9 Mon Sep 17 00:00:00 2001 From: stuppie Date: Thu, 3 Sep 2026 16:54:17 -0600 Subject: now a lot of thl wall/session tests working --- test_utils/models/thl/conftest.py | 63 +++++++++++++++--------------------- tests/models/thl/test_adjustments.py | 9 ++---- 2 files changed, 28 insertions(+), 44 deletions(-) (limited to 'test_utils/models') diff --git a/test_utils/models/thl/conftest.py b/test_utils/models/thl/conftest.py index 6f2835e..80a80e3 100644 --- a/test_utils/models/thl/conftest.py +++ b/test_utils/models/thl/conftest.py @@ -4,7 +4,6 @@ from collections.abc import Callable from datetime import UTC, datetime, timedelta from decimal import ROUND_DOWN, Decimal from random import choice as rand_choice -from random import choice as randchoice from random import randint, random from typing import TYPE_CHECKING, Any from uuid import uuid4 @@ -14,29 +13,26 @@ import pytest from grip_client.enums import AccessType from pydantic import PositiveInt +from generalresearch.currency import USDCent from generalresearch.managers.thl.payout import UserPayoutEventManager from generalresearch.models.custom_types import ( AwareDatetimeISO, IPvAnyAddressStr, UUIDStr, ) +from generalresearch.models.definitions import DeviceType, Source 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.payout import UserPayoutEvent 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 from generalresearch.pg_helper import PostgresConfig if TYPE_CHECKING: - from generalresearch.currency import USDCent from generalresearch.managers.thl.ipinfo import ( IPGeonameManager, IPInformationManager, @@ -51,7 +47,6 @@ if TYPE_CHECKING: 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 @@ -94,7 +89,7 @@ fake = faker.Faker() @pytest.fixture def wall_factory( wall_manager: WallManager, - session_factory: Callable[..., Session], + bare_session_factory: Callable[..., Session], session_manager: SessionManager, ) -> Callable[..., Wall]: @@ -143,7 +138,7 @@ def wall_factory( session_manager.get_from_id(session_id=session_id) if session_id else None - ) or session_factory(save=True, user_id=user_id) + ) or bare_session_factory(save=True, user_id=user_id) assert session, "Wall factory requires Session" @@ -205,7 +200,10 @@ def wall_status() -> Status: @pytest.fixture -def session_factory(session_manager: SessionManager, user_factory: Callable[..., User]): +def bare_session_factory( + session_manager: SessionManager, user_factory: Callable[..., User] +): + # Create a session with no wall events def _inner( save: bool = True, @@ -251,40 +249,33 @@ def session_factory(session_manager: SessionManager, user_factory: Callable[..., @pytest.fixture() -def session(session_factory: Callable[..., Session]) -> Session: - return session_factory(save=True) +def bare_session(bare_session_factory: Callable[..., Session]) -> Session: + # A session with no wall events + return bare_session_factory() @pytest.fixture -def session_w_wall( - user: User, - session: Session, +def session( + bare_session: Session, wall_factory: Callable[..., Wall], ) -> Session: - + s = bare_session.model_copy() wall: Wall = wall_factory( - session_id=session.id, - user_id=session.user_id, - started=session.started, + session_id=s.id, + user_id=s.user_id, + started=s.started, ) - session.append_wall_event(w=wall) - - return session - - -@pytest.fixture() -def unsaved_session(session_factory: Callable[..., Session]) -> Session: - return session_factory(save=False) + s.append_wall_event(w=wall) + return s @pytest.fixture -def session_w_wall_factory( +def session_factory( wall_manager: WallManager, utc_hour_ago: datetime, - session_factory: Callable[..., Session], + bare_session_factory: Callable[..., Session], wall_factory: Callable[..., Wall], ) -> Callable[..., Session]: - from generalresearch.models.thl.session import Source def _inner( user: User, @@ -303,7 +294,7 @@ def session_w_wall_factory( if wall_statuses: assert len(wall_statuses) == wall_count - s = session_factory(started=started, user=user, country_iso="us") + s = bare_session_factory(started=started, user=user, country_iso="us") for idx in range(wall_count): if idx == 0: # First Wall Event in a session @@ -334,7 +325,7 @@ def session_w_wall_factory( wall_manager.finish( wall=w, status=_final_status, - status_code_1=randchoice(options), + status_code_1=rand_choice(options), finished=w.started + timedelta(seconds=randint(a=60 * 2, b=60 * 10)), ) @@ -345,11 +336,10 @@ def session_w_wall_factory( @pytest.fixture(scope="function") def finished_session_factory( - session_w_wall_factory: Callable[..., Session], + session_factory: Callable[..., Session], session_manager: SessionManager, utc_hour_ago: datetime, ) -> Callable[..., Session]: - from generalresearch.models.thl.session import Source def _inner( user: User, @@ -363,7 +353,7 @@ def finished_session_factory( final_status: Status = Status.COMPLETE, started: datetime = utc_hour_ago, ) -> Session: - s: Session = session_w_wall_factory( + s: Session = session_factory( user=user, wall_count=wall_count, wall_req_cpi=wall_req_cpi, @@ -804,7 +794,6 @@ def brokerage_product_payout_event_factory( 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)) diff --git a/tests/models/thl/test_adjustments.py b/tests/models/thl/test_adjustments.py index cd75318..5d8605a 100644 --- a/tests/models/thl/test_adjustments.py +++ b/tests/models/thl/test_adjustments.py @@ -14,15 +14,13 @@ from generalresearch.models.thl.session import ( Status, StatusCode1, WallAdjustedStatus, + Session, + Wall, ) 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) @@ -36,7 +34,6 @@ adj_ts3 = datetime(2023, 2, 4, tzinfo=UTC) class TestProductAdjustments: - @pytest.mark.parametrize("payout", [".6", "1", "1.8", "2", "500.0000"]) def test_determine_bp_payment_no_rounding( self, product_factory: Callable[..., Product], payout: str @@ -57,7 +54,6 @@ class TestProductAdjustments: class TestSessionAdjustments: - def test_status_complete(self, session_factory: Callable[..., Session], user: User): # Completed Session with 2 wall events s1 = session_factory( @@ -80,7 +76,6 @@ class TestSessionAdjustments: class TestAdjustments: - def test_finish_with_status( self, session_factory: Callable[..., Session], -- cgit v1.2.3 From c9e1fc6839d8fc3ce145f65ad3c3c628af67626b Mon Sep 17 00:00:00 2001 From: stuppie Date: Thu, 3 Sep 2026 18:03:41 -0600 Subject: working TestThlLedgerTxManager. working wall manager --- test_utils/models/conftest.py | 5 +- test_utils/models/thl/conftest.py | 22 ++--- tests/managers/thl/test_ledger/test_thl_lm_tx.py | 108 +++++++---------------- tests/managers/thl/test_task_adjustment.py | 2 +- 4 files changed, 39 insertions(+), 98 deletions(-) (limited to 'test_utils/models') diff --git a/test_utils/models/conftest.py b/test_utils/models/conftest.py index 43f18c1..ffce272 100644 --- a/test_utils/models/conftest.py +++ b/test_utils/models/conftest.py @@ -12,7 +12,6 @@ from generalresearch.models.thl.survey.model import Buyer, Survey if TYPE_CHECKING: from generalresearch.managers.thl.buyer import BuyerManager - from generalresearch.managers.thl.product import ProductManager from generalresearch.managers.thl.survey import SurveyManager from generalresearch.models.thl.product import ( PayoutConfig, @@ -47,7 +46,6 @@ def payout_config(request: Request) -> PayoutConfig: def product_user_wallet_yes( product_factory: Callable[..., Product], payout_config: PayoutConfig, - product_manager: ProductManager, ) -> Product: from generalresearch.models.thl.product import UserWalletConfig @@ -58,7 +56,7 @@ def product_user_wallet_yes( @pytest.fixture def product_user_wallet_no( - product_factory: Callable[..., Product], product_manager: ProductManager + product_factory: Callable[..., Product], ) -> Product: from generalresearch.models.thl.product import UserWalletConfig @@ -68,7 +66,6 @@ def product_user_wallet_no( @pytest.fixture def product_amt_true( product_factory: Callable[..., Product], - product_manager: ProductManager, payout_config: PayoutConfig, ) -> Product: from generalresearch.models.thl.product import UserWalletConfig diff --git a/test_utils/models/thl/conftest.py b/test_utils/models/thl/conftest.py index 80a80e3..9cb68be 100644 --- a/test_utils/models/thl/conftest.py +++ b/test_utils/models/thl/conftest.py @@ -76,21 +76,12 @@ fake = faker.Faker() # --- Wall --- -# 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, bare_session_factory: Callable[..., Session], session_manager: SessionManager, + user_factory: Callable[..., User], ) -> Callable[..., Wall]: def _inner( @@ -98,7 +89,7 @@ def wall_factory( save: bool = True, session: Session | None = None, session_id: PositiveInt | None = None, - user_id: int | None = None, + user: User | None = None, started: datetime | None = None, source: Source | None = None, req_survey_id: str | None = None, @@ -107,9 +98,8 @@ def wall_factory( uuid_id: str | None = None, ) -> Wall: """To be used in tests, where we don't care about certain fields""" - + user = user or user_factory() if save: - 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), @@ -138,7 +128,7 @@ def wall_factory( session_manager.get_from_id(session_id=session_id) if session_id else None - ) or bare_session_factory(save=True, user_id=user_id) + ) or bare_session_factory(save=True, user=user) assert session, "Wall factory requires Session" @@ -262,7 +252,7 @@ def session( s = bare_session.model_copy() wall: Wall = wall_factory( session_id=s.id, - user_id=s.user_id, + user=s.user, started=s.started, ) s.append_wall_event(w=wall) @@ -308,7 +298,7 @@ def session_factory( w = wall_factory( session_id=s.id, source=wall_source, - user_id=s.user_id, + user=s.user, started=wall_started, req_cpi=wall_req_cpis[idx] if wall_req_cpis else wall_req_cpi, ) 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 cda88da..aa3b378 100644 --- a/tests/managers/thl/test_ledger/test_thl_lm_tx.py +++ b/tests/managers/thl/test_ledger/test_thl_lm_tx.py @@ -11,6 +11,9 @@ from uuid import uuid4 import pytest from generalresearch.currency import USDCent +from generalresearch.managers.thl.ledger_manager.exceptions import ( + LedgerTransactionConditionFailedError, +) from generalresearch.managers.thl.ledger_manager.ledger import ( LedgerTransaction, ) @@ -56,17 +59,19 @@ logger = logging.getLogger("LedgerManager") class TestThlLedgerTxManager: + @pytest.fixture(autouse=True) + def setup(self, delete_ledger_db, create_main_accounts): + delete_ledger_db() + create_main_accounts() def test_create_tx_task_complete( self, 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_ledger_manager.create_tx_task_complete(wall=wall, user=user) assert isinstance(tx, LedgerTransaction) @@ -91,14 +96,11 @@ class TestThlLedgerTxManager: self, 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_code_1 = s1.determine_session_status() @@ -123,15 +125,12 @@ class TestThlLedgerTxManager: 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, product_factory: Callable[..., Product], ): - delete_ledger_db() - create_main_accounts() + product = product_factory( payout_config=PayoutConfig( payout_transformation=PayoutTransformation( @@ -168,7 +167,6 @@ class TestThlLedgerTxManager: self, session_factory: Callable[..., Session], user: User, - create_main_accounts: Callable[..., None], thl_ledger_manager: ThlLedgerManager, ledger_manager: LedgerManager, session_manager: SessionManager, @@ -196,9 +194,8 @@ class TestThlLedgerTxManager: def test_create_tx_task_adjustment( self, wall_factory: Callable[..., Wall], - session: Session, + bare_session: Session, user: User, - create_main_accounts: Callable[..., None], thl_ledger_manager: ThlLedgerManager, ledger_manager: LedgerManager, ): @@ -208,9 +205,8 @@ class TestThlLedgerTxManager: the transaction comes back with balanced amounts, and that the name of the Source is in the Tx description """ - wall_status = Status.COMPLETE - wall: Wall = wall_factory(session=session, wall_status=wall_status) + wall: Wall = wall_factory(session=bare_session, wall_status=wall_status) tx = thl_ledger_manager.create_tx_task_adjustment(wall=wall, user=user) assert isinstance(tx, LedgerTransaction) @@ -232,12 +228,11 @@ class TestThlLedgerTxManager: status, status_code_1 = session.determine_session_status() thl_net, commission_amount, bp_pay, user_pay = session.determine_payments() - # The default session fixture is just an unfinished wall event assert len(session.wall_events) == 1 assert session.finished is None - assert status == Status.TIMEOUT + assert status == Status.FAIL assert status_code_1 in list( - WALL_ALLOWED_STATUS_STATUS_CODE.get(Status.TIMEOUT, {}) + WALL_ALLOWED_STATUS_STATUS_CODE.get(Status.FAIL, {}) ) assert thl_net == Decimal(0) assert commission_amount == Decimal(0) @@ -246,7 +241,9 @@ class TestThlLedgerTxManager: # Update the finished timestamp, but nothing else. This means that # there is no financial changes needed - session.update(finished=datetime.now(tz=UTC) + timedelta(minutes=10)) + session.update( + finished=datetime.now(tz=UTC) + timedelta(minutes=10), status=Status.FAIL + ) assert session.finished with caplog.at_level(logging.INFO): tx = thl_ledger_manager.create_tx_bp_adjustment(session=session) @@ -295,8 +292,9 @@ class TestThlLedgerTxManager: assert balance == int(rand_amount) * -1 # Test some basic assertions - with caplog.at_level(logging.INFO), pytest.raises( - expected_exception=ValueError + with ( + caplog.at_level(logging.INFO), + pytest.raises(expected_exception=LedgerTransactionConditionFailedError), ): thl_ledger_manager.create_tx_bp_payout( product=product, @@ -339,7 +337,6 @@ class TestThlLedgerTxManager: def test_create_tx_plug_bp_wallet( self, product: Product, - create_main_accounts: Callable[..., None], thl_ledger_manager: ThlLedgerManager, ledger_manager: LedgerManager, currency: LedgerCurrency, @@ -369,7 +366,6 @@ class TestThlLedgerTxManager: def test_create_tx_plug_bp_wallet_( self, product: Product, - create_main_accounts: Callable[..., None], thl_ledger_manager: ThlLedgerManager, ledger_manager: LedgerManager, currency: LedgerCurrency, @@ -478,11 +474,9 @@ class TestThlLedgerTxManager: 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() pe = UserPayoutEvent( uuid=uuid4().hex, @@ -508,14 +502,10 @@ class TestThlLedgerTxManager: self, 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_ledger_manager.get_account_or_create_user_wallet(user=user) @@ -575,7 +565,6 @@ class TestThlLedgerTxManager: self, user_factory: Callable[..., User], product_user_wallet_yes: Product, - create_main_accounts: Callable[..., None], thl_ledger_manager: ThlLedgerManager, ledger_manager: LedgerManager, ): @@ -626,7 +615,6 @@ class TestThlLedgerTxManager: self, user_factory: Callable[..., User], product_user_wallet_yes: Product, - create_main_accounts: Callable[..., None], thl_ledger_manager: ThlLedgerManager, ledger_manager: LedgerManager, currency: LedgerCurrency, @@ -673,7 +661,6 @@ class TestThlLedgerTxManager: self, user_factory: Callable[..., User], product_user_wallet_yes: Product, - create_main_accounts: Callable[..., None], thl_ledger_manager: ThlLedgerManager, ledger_manager: LedgerManager, currency: LedgerCurrency, @@ -715,7 +702,6 @@ class TestThlLedgerTxManager: self, user_factory: Callable[..., User], product_user_wallet_yes: Product, - create_main_accounts: Callable[..., None], thl_ledger_manager: ThlLedgerManager, ledger_manager: LedgerManager, currency: LedgerCurrency, @@ -750,7 +736,6 @@ class TestThlLedgerTxManager: self, user_factory: Callable[..., User], product_user_wallet_yes: Product, - create_main_accounts: Callable[..., None], thl_ledger_manager: ThlLedgerManager, ledger_manager: LedgerManager, currency: LedgerCurrency, @@ -786,17 +771,18 @@ class TestThlLedgerTxManagerFlows: examples """ + @pytest.fixture(autouse=True) + def setup(self, delete_ledger_db, create_main_accounts): + delete_ledger_db() + create_main_accounts() + def test_create_tx_task_complete( 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() wall1 = Wall( user_id=1, @@ -868,7 +854,6 @@ class TestThlLedgerTxManagerFlows: def test_create_transaction_task_complete_1_cent( self, user: User, - create_main_accounts: Callable[..., None], thl_ledger_manager: ThlLedgerManager, ledger_manager: LedgerManager, currency: LedgerCurrency, @@ -893,16 +878,12 @@ class TestThlLedgerTxManagerFlows: def test_create_transaction_bp_payment( self, 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() s1: Session = session_factory( user=user, @@ -956,7 +937,6 @@ class TestThlLedgerTxManagerFlows: self, user_factory: Callable[..., User], product_user_wallet_no: Product, - create_main_accounts: Callable[..., None], thl_ledger_manager: ThlLedgerManager, ledger_manager: LedgerManager, currency: LedgerCurrency, @@ -1000,15 +980,12 @@ class TestThlLedgerTxManagerFlows: def test_create_transaction_bp_payment_round2( 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() + # user must be no user wallet # e.g. session 869b5bfa47f44b4f81cd095ed01df2ff this fails if you dont round properly @@ -1044,7 +1021,6 @@ class TestThlLedgerTxManagerFlows: self, user_factory: Callable[..., User], product_user_wallet_yes: Product, - create_main_accounts: Callable[..., None], thl_ledger_manager: ThlLedgerManager, ledger_manager: LedgerManager, currency: LedgerCurrency, @@ -1086,8 +1062,6 @@ class TestThlLedgerTxManagerFlows: self, 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, @@ -1096,8 +1070,6 @@ class TestThlLedgerTxManagerFlows: 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) @@ -1165,20 +1137,20 @@ class TestThlLedgerTxManagerFlows: class TestThlLedgerManagerAdj: + @pytest.fixture(autouse=True) + def setup(self, delete_ledger_db, create_main_accounts): + delete_ledger_db() + create_main_accounts() def test_create_tx_task_adjustment( self, 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() user: User = user_factory(product=product_user_wallet_no) @@ -1272,7 +1244,6 @@ class TestThlLedgerManagerAdj: self, user: User, product_user_wallet_no: Product, - create_main_accounts: Callable[..., None], caplog, thl_ledger_manager: ThlLedgerManager, ledger_manager: LedgerManager, @@ -1281,10 +1252,7 @@ class TestThlLedgerManagerAdj: wall_manager: WallManager, session_factory: Callable[..., Session], utc_hour_ago: datetime, - delete_ledger_db: Callable[..., None], ): - delete_ledger_db() - create_main_accounts() s1 = session_factory( user=user, @@ -1391,15 +1359,11 @@ class TestThlLedgerManagerAdj: self, 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() # This failed when I didn't check that `change_commission` > 0 in # create_transaction_bp_adjustment @@ -1447,9 +1411,7 @@ class TestThlLedgerManagerAdj: self, 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_ledger_manager: ThlLedgerManager, ledger_manager: LedgerManager, @@ -1458,8 +1420,7 @@ class TestThlLedgerManagerAdj: session_manager: SessionManager, wall_manager: WallManager, ): - delete_ledger_db() - create_main_accounts() + user: User = user_factory(product=product_user_wallet_no) s1: Session = session_factory( user=user, final_status=Status.ABANDON, wall_req_cpi=Decimal(1) @@ -1521,15 +1482,11 @@ class TestThlLedgerManagerAdj: self, user_factory: Callable[..., User], product_user_wallet_yes: Product, - create_main_accounts: Callable[..., None], - delete_ledger_db: Callable[..., None], caplog, thl_ledger_manager: ThlLedgerManager, ledger_manager: LedgerManager, currency: LedgerCurrency, ): - delete_ledger_db() - create_main_accounts() now = datetime.now(UTC) - timedelta(days=1) user: User = user_factory(product=product_user_wallet_yes) @@ -1746,16 +1703,13 @@ class TestThlLedgerManagerAdj: self, user_factory: Callable[..., User], product_user_wallet_no: Product, - create_main_accounts: Callable[..., None], - delete_ledger_db: Callable[..., None], caplog, thl_ledger_manager: ThlLedgerManager, ledger_manager: LedgerManager, utc_hour_ago: datetime, currency: LedgerCurrency, ): - delete_ledger_db() - create_main_accounts() + user: User = user_factory(product=product_user_wallet_no) wall1 = Wall( diff --git a/tests/managers/thl/test_task_adjustment.py b/tests/managers/thl/test_task_adjustment.py index 323d6db..11dbec4 100644 --- a/tests/managers/thl/test_task_adjustment.py +++ b/tests/managers/thl/test_task_adjustment.py @@ -56,7 +56,7 @@ def session_fail( session = session_factory(started=datetime.now(UTC), user=user) wall1 = wall_factory( session_id=session.id, - user_id=user.user_id, + user=user, source=Source.DYNATA, req_survey_id="72723", req_cpi=Decimal("3.22"), -- cgit v1.2.3 From 4b705346968e38671ac9601cfa8444c944ecc8bd Mon Sep 17 00:00:00 2001 From: stuppie Date: Fri, 4 Sep 2026 12:52:46 -0600 Subject: remove a lot of if type_checking; fix model_rebuild errors. Remove all network models/managers/test/fixtures --- generalresearch/managers/network/__init__.py | 0 generalresearch/managers/network/label.py | 152 ------- generalresearch/managers/network/mtr.py | 53 --- generalresearch/managers/network/nmap.py | 60 --- generalresearch/managers/network/rdns.py | 37 -- generalresearch/managers/network/tool_run.py | 138 ------- generalresearch/models/cint/question.py | 12 +- generalresearch/models/cint/survey.py | 13 +- generalresearch/models/cint/task_collection.py | 6 +- generalresearch/models/dynata/survey.py | 10 +- generalresearch/models/dynata/task_collection.py | 6 +- generalresearch/models/innovate/question.py | 12 +- generalresearch/models/innovate/survey.py | 10 +- generalresearch/models/innovate/task_collection.py | 6 +- generalresearch/models/legacy/offerwall.py | 37 +- generalresearch/models/legacy/questions.py | 14 +- generalresearch/models/lucid/question.py | 12 +- generalresearch/models/lucid/survey.py | 6 +- generalresearch/models/morning/question.py | 6 +- generalresearch/models/morning/survey.py | 19 +- generalresearch/models/morning/task_collection.py | 6 +- generalresearch/models/network/__init__.py | 0 generalresearch/models/network/definitions.py | 70 ---- generalresearch/models/network/label.py | 129 ------ generalresearch/models/network/mtr/__init__.py | 0 generalresearch/models/network/mtr/command.py | 76 ---- generalresearch/models/network/mtr/execute.py | 60 --- generalresearch/models/network/mtr/parser.py | 19 - generalresearch/models/network/mtr/result.py | 175 --------- generalresearch/models/network/nmap/__init__.py | 0 generalresearch/models/network/nmap/command.py | 51 --- generalresearch/models/network/nmap/execute.py | 56 --- generalresearch/models/network/nmap/parser.py | 414 ------------------- generalresearch/models/network/nmap/result.py | 436 --------------------- generalresearch/models/network/rdns/__init__.py | 0 generalresearch/models/network/rdns/command.py | 38 -- generalresearch/models/network/rdns/execute.py | 44 --- generalresearch/models/network/rdns/parser.py | 21 - generalresearch/models/network/rdns/result.py | 52 --- generalresearch/models/network/tool_run.py | 121 ------ generalresearch/models/network/tool_run_command.py | 68 ---- generalresearch/models/network/utils.py | 5 - generalresearch/models/pollfish/question.py | 10 +- generalresearch/models/precision/question.py | 12 +- generalresearch/models/precision/survey.py | 7 +- .../models/precision/task_collection.py | 6 +- generalresearch/models/prodege/question.py | 14 +- generalresearch/models/prodege/survey.py | 11 +- generalresearch/models/prodege/task_collection.py | 6 +- generalresearch/models/repdata/question.py | 10 +- generalresearch/models/repdata/task_collection.py | 6 +- generalresearch/models/sago/question.py | 12 +- generalresearch/models/sago/survey.py | 20 +- generalresearch/models/sago/task_collection.py | 6 +- generalresearch/models/spectrum/question.py | 14 +- generalresearch/models/spectrum/survey.py | 4 +- generalresearch/models/spectrum/task_collection.py | 6 +- generalresearch/models/thl/contest/__init__.py | 14 +- generalresearch/models/thl/contest/contest.py | 10 +- .../models/thl/contest/contest_entry.py | 6 +- generalresearch/models/thl/contest/leaderboard.py | 40 +- generalresearch/models/thl/contest/raffle.py | 20 +- .../models/thl/profiling/upk_property.py | 5 +- generalresearch/models/thl/profiling/user_info.py | 14 +- generalresearch/models/thl/survey/__init__.py | 21 +- test_utils/managers/network/__init__.py | 0 test_utils/managers/network/conftest.py | 0 test_utils/models/network/__init__.py | 0 test_utils/models/network/conftest.py | 145 ------- tests/managers/network/__init__.py | 0 tests/managers/network/test_label.py | 209 ---------- tests/managers/network/test_tool_run.py | 25 -- tests/models/network/__init__.py | 0 tests/models/network/test_mtr.py | 33 -- tests/models/network/test_nmap.py | 39 -- tests/models/network/test_nmap_parser.py | 32 -- tests/models/network/test_rdns.py | 40 -- 77 files changed, 169 insertions(+), 3078 deletions(-) delete mode 100644 generalresearch/managers/network/__init__.py delete mode 100644 generalresearch/managers/network/label.py delete mode 100644 generalresearch/managers/network/mtr.py delete mode 100644 generalresearch/managers/network/nmap.py delete mode 100644 generalresearch/managers/network/rdns.py delete mode 100644 generalresearch/managers/network/tool_run.py delete mode 100644 generalresearch/models/network/__init__.py delete mode 100644 generalresearch/models/network/definitions.py delete mode 100644 generalresearch/models/network/label.py delete mode 100644 generalresearch/models/network/mtr/__init__.py delete mode 100644 generalresearch/models/network/mtr/command.py delete mode 100644 generalresearch/models/network/mtr/execute.py delete mode 100644 generalresearch/models/network/mtr/parser.py delete mode 100644 generalresearch/models/network/mtr/result.py delete mode 100644 generalresearch/models/network/nmap/__init__.py delete mode 100644 generalresearch/models/network/nmap/command.py delete mode 100644 generalresearch/models/network/nmap/execute.py delete mode 100644 generalresearch/models/network/nmap/parser.py delete mode 100644 generalresearch/models/network/nmap/result.py delete mode 100644 generalresearch/models/network/rdns/__init__.py delete mode 100644 generalresearch/models/network/rdns/command.py delete mode 100644 generalresearch/models/network/rdns/execute.py delete mode 100644 generalresearch/models/network/rdns/parser.py delete mode 100644 generalresearch/models/network/rdns/result.py delete mode 100644 generalresearch/models/network/tool_run.py delete mode 100644 generalresearch/models/network/tool_run_command.py delete mode 100644 generalresearch/models/network/utils.py delete mode 100644 test_utils/managers/network/__init__.py delete mode 100644 test_utils/managers/network/conftest.py delete mode 100644 test_utils/models/network/__init__.py delete mode 100644 test_utils/models/network/conftest.py delete mode 100644 tests/managers/network/__init__.py delete mode 100644 tests/managers/network/test_label.py delete mode 100644 tests/managers/network/test_tool_run.py delete mode 100644 tests/models/network/__init__.py delete mode 100644 tests/models/network/test_mtr.py delete mode 100644 tests/models/network/test_nmap.py delete mode 100644 tests/models/network/test_nmap_parser.py delete mode 100644 tests/models/network/test_rdns.py (limited to 'test_utils/models') diff --git a/generalresearch/managers/network/__init__.py b/generalresearch/managers/network/__init__.py deleted file mode 100644 index e69de29..0000000 diff --git a/generalresearch/managers/network/label.py b/generalresearch/managers/network/label.py deleted file mode 100644 index cdea016..0000000 --- a/generalresearch/managers/network/label.py +++ /dev/null @@ -1,152 +0,0 @@ -from __future__ import annotations - -from collections.abc import Collection -from datetime import UTC, datetime, timedelta -from typing import TYPE_CHECKING - -from psycopg import sql -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 - -if TYPE_CHECKING: - - from generalresearch.models.network.label import IPLabelKind, IPLabelSource - - -class IPLabelManager(PostgresManager): - def create(self, ip_label: IPLabel) -> IPLabel: - query = sql.SQL(""" - INSERT INTO network_iplabel ( - ip, labeled_at, created_at, - label_kind, source, confidence, - provider, metadata - ) VALUES ( - %(ip)s, %(labeled_at)s, %(created_at)s, - %(label_kind)s, %(source)s, %(confidence)s, - %(provider)s, %(metadata)s - ) RETURNING id;""") - params = ip_label.model_dump_postgres() - with self.pg_config.make_connection() as conn, conn.cursor() as c: - c.execute(query, params) - _pk = c.fetchone()["id"] - return ip_label - - def make_filter_str( - self, - ips: Collection[IPvAnyNetworkStr] | None = None, - ip_in_network: IPvAnyAddressStr | None = None, - label_kind: IPLabelKind | None = None, - source: IPLabelSource | None = None, - labeled_at: AwareDatetimeISO | None = None, - labeled_after: AwareDatetimeISO | None = None, - labeled_before: AwareDatetimeISO | None = None, - provider: str | None = None, - ): - filters = [] - params = {} - if labeled_after or labeled_before: - 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") - params["time_start"] = time_start - params["time_end"] = time_end - if labeled_at: - assert labeled_at.tzinfo.utcoffset(labeled_at) == timedelta(), "must be UTC" - filters.append("labeled_at == %(labeled_at)s") - params["labeled_at"] = labeled_at - if label_kind: - filters.append("label_kind = %(label_kind)s") - params["label_kind"] = label_kind.value - if source: - filters.append("source = %(source)s") - params["source"] = source.value - if provider: - filters.append("provider = %(provider)s") - params["provider"] = provider - if ips is not None: - filters.append("ip = ANY(%(ips)s)") - params["ips"] = list(ips) - if ip_in_network: - """ - Return matching networks. - e.g. ip = '13f9:c462:e039:a38c::1', might return rows - where ip = '13f9:c462:e039::/48' or '13f9:c462:e039:a38c::/64' - """ - filters.append("ip >>= %(ip_in_network)s") - params["ip_in_network"] = ip_in_network - - filter_str = "WHERE " + " AND ".join(filters) if filters else "" - return filter_str, params - - def filter( - self, - ips: Collection[IPvAnyNetworkStr] | None = None, - ip_in_network: IPvAnyAddressStr | None = None, - label_kind: IPLabelKind | None = None, - source: IPLabelSource | None = None, - labeled_at: AwareDatetimeISO | None = None, - labeled_after: AwareDatetimeISO | None = None, - labeled_before: AwareDatetimeISO | None = None, - provider: str | None = None, - ) -> list[IPLabel]: - filter_str, params = self.make_filter_str( - ips=ips, - ip_in_network=ip_in_network, - label_kind=label_kind, - source=source, - labeled_at=labeled_at, - labeled_after=labeled_after, - labeled_before=labeled_before, - provider=provider, - ) - query = f""" - SELECT - ip, labeled_at, created_at, - label_kind, source, confidence, - provider, metadata - FROM network_iplabel - {filter_str} - """ - res = self.pg_config.execute_sql_query(query, params) - return [IPLabel.model_validate(rec) for rec in res] - - def get_most_specific_matching_network(self, ip: IPvAnyAddressStr) -> IPvAnyNetwork: - """ - e.g. ip = 'b5f4:dc2:f136:70d5:5b6e:9a85:c7d4:3517', might return - 'b5f4:dc2:f136:70d5::/64' - """ - ip = TypeAdapter(IPvAnyAddressStr).validate_python(ip) - - query = """ - SELECT ip - FROM network_iplabel - WHERE ip >>= %(ip)s - ORDER BY masklen(ip) DESC - LIMIT 1;""" - res = self.pg_config.execute_sql_query(query, {"ip": ip}) - if res: - return IPvAnyNetwork(res[0]["ip"]) - - def test_join(self, ip): - query = """ - SELECT - to_jsonb(i) AS ipinfo, - to_jsonb(l) AS iplabel - FROM thl_ipinformation i - LEFT JOIN network_iplabel l - ON l.ip >>= i.ip - WHERE i.ip = %(ip)s - ORDER BY masklen(l.ip) DESC;""" - params = {"ip": ip} - res = self.pg_config.execute_sql_query(query, params) - return res diff --git a/generalresearch/managers/network/mtr.py b/generalresearch/managers/network/mtr.py deleted file mode 100644 index 7b79d96..0000000 --- a/generalresearch/managers/network/mtr.py +++ /dev/null @@ -1,53 +0,0 @@ -from __future__ import annotations - -from typing import TYPE_CHECKING - -from psycopg import Cursor, sql - -from generalresearch.managers.base import PostgresManager - -if TYPE_CHECKING: - from generalresearch.models.network.tool_run import MTRRun - - -class MTRRunManager(PostgresManager): - - def _create(self, run: MTRRun, c: Cursor | None = None) -> None: - """ - Do not use this directly. Must only be used in the context of a toolrun - """ - query = sql.SQL(""" - INSERT INTO network_mtr ( - run_id, source_ip, facility_id, - protocol, port, parsed, - started_at, ip, scan_group_id - ) - VALUES ( - %(run_id)s, %(source_ip)s, %(facility_id)s, - %(protocol)s, %(port)s, %(parsed)s, - %(started_at)s, %(ip)s, %(scan_group_id)s - ); - """) - params = run.model_dump_postgres() - - query_hops = sql.SQL(""" - INSERT INTO network_mtrhop ( - hop, ip, domain, asn, mtr_run_id - ) VALUES ( - %(hop)s, %(ip)s, %(domain)s, - %(asn)s, %(mtr_run_id)s - ) - """) - mtr_run = run.parsed - params_hops = [h.model_dump_postgres(run_id=run.id) for h in mtr_run.hops] - - if c: - c.execute(query, params) - if params_hops: - c.executemany(query_hops, params_hops) - - else: - with self.pg_config.make_connection() as conn, conn.cursor() as _c: - _c.execute(query, params) - if params_hops: - _c.executemany(query_hops, params_hops) diff --git a/generalresearch/managers/network/nmap.py b/generalresearch/managers/network/nmap.py deleted file mode 100644 index 96c6009..0000000 --- a/generalresearch/managers/network/nmap.py +++ /dev/null @@ -1,60 +0,0 @@ -from __future__ import annotations - -from typing import TYPE_CHECKING - -from psycopg import Cursor, sql - -from generalresearch.managers.base import PostgresManager - -if TYPE_CHECKING: - from generalresearch.models.network.tool_run import NmapRun - - -class NmapRunManager(PostgresManager): - - def _create(self, run: NmapRun, c: Cursor | None = None) -> None: - """ - Insert a PortScan + PortScanPorts from a Pydantic NmapResult. - Do not use this directly. Must only be used in the context of a toolrun - """ - query = sql.SQL(""" - INSERT INTO network_portscan ( - run_id, xml_version, host_state, - host_state_reason, latency_ms, distance, - uptime_seconds, last_boot, - parsed, scan_group_id, open_tcp_ports, - started_at, ip, open_udp_ports - ) - VALUES ( - %(run_id)s, %(xml_version)s, %(host_state)s, - %(host_state_reason)s, %(latency_ms)s, %(distance)s, - %(uptime_seconds)s, %(last_boot)s, - %(parsed)s, %(scan_group_id)s, %(open_tcp_ports)s, - %(started_at)s, %(ip)s, %(open_udp_ports)s - ); - """) - params = run.model_dump_postgres() - - query_ports = sql.SQL(""" - INSERT INTO network_portscanport ( - port_scan_id, protocol, port, - state, reason, reason_ttl, - service_name - ) VALUES ( - %(port_scan_id)s, %(protocol)s, %(port)s, - %(state)s, %(reason)s, %(reason_ttl)s, - %(service_name)s - ) - """) - nmap_run = run.parsed - params_ports = [p.model_dump_postgres(run_id=run.id) for p in nmap_run.ports] - - if c: - c.execute(query, params) - if nmap_run.ports: - c.executemany(query_ports, params_ports) - else: - with self.pg_config.make_connection() as conn, conn.cursor(): - c.execute(query, params) - if nmap_run.ports: - c.executemany(query_ports, params_ports) diff --git a/generalresearch/managers/network/rdns.py b/generalresearch/managers/network/rdns.py deleted file mode 100644 index 1800364..0000000 --- a/generalresearch/managers/network/rdns.py +++ /dev/null @@ -1,37 +0,0 @@ -from __future__ import annotations - -from typing import TYPE_CHECKING - -from psycopg import Cursor - -from generalresearch.managers.base import PostgresManager - -if TYPE_CHECKING: - from generalresearch.models.network.tool_run import RDNSRun - - -class RDNSRunManager(PostgresManager): - - def _create(self, run: RDNSRun, c: Cursor | None = None) -> None: - """ - Do not use this directly. Must only be used in the context of a toolrun - """ - query = """ - INSERT INTO network_rdnsresult ( - run_id, primary_hostname, primary_domain, - hostname_count, hostnames, - ip, started_at, scan_group_id - ) - VALUES ( - %(run_id)s, %(primary_hostname)s, %(primary_domain)s, - %(hostname_count)s, %(hostnames)s, - %(ip)s, %(started_at)s, %(scan_group_id)s - ); - """ - params = run.model_dump_postgres() - if c: - c.execute(query, params) - - else: - with self.pg_config.make_connection() as conn, conn.cursor() as _c: - _c.execute(query, params) diff --git a/generalresearch/managers/network/tool_run.py b/generalresearch/managers/network/tool_run.py deleted file mode 100644 index ec06305..0000000 --- a/generalresearch/managers/network/tool_run.py +++ /dev/null @@ -1,138 +0,0 @@ -from __future__ import annotations - -from collections.abc import Collection -from typing import TYPE_CHECKING - -from psycopg import Cursor, sql - -from generalresearch.managers.base import PostgresManager -from generalresearch.managers.network.mtr import MTRRunManager -from generalresearch.managers.network.nmap import NmapRunManager -from generalresearch.managers.network.rdns import RDNSRunManager -from generalresearch.models.network.rdns.result import RDNSResult -from generalresearch.models.network.tool_run import ( - MTRRun, - NmapRun, - RDNSRun, - ToolRun, -) - -if TYPE_CHECKING: - from generalresearch.managers.base import Permission - from generalresearch.models.network.tool_run import ToolName - from generalresearch.pg_helper import PostgresConfig - - -class ToolRunManager(PostgresManager): - def __init__( - self, - pg_config: PostgresConfig, - permissions: Collection[Permission] | None = None, - ): - super().__init__(pg_config=pg_config, permissions=permissions) - self.nmap_manager = NmapRunManager(self.pg_config) - self.rdns_manager = RDNSRunManager(self.pg_config) - self.mtr_manager = MTRRunManager(self.pg_config) - - def _create_tool_run(self, run: NmapRun | RDNSRun | MTRRun, c: Cursor): - query = sql.SQL(""" - INSERT INTO network_toolrun ( - ip, scan_group_id, tool_class, - tool_name, tool_version, started_at, - finished_at, status, raw_command, - config - ) - VALUES ( - %(ip)s, %(scan_group_id)s, %(tool_class)s, - %(tool_name)s, %(tool_version)s, %(started_at)s, - %(finished_at)s, %(status)s, %(raw_command)s, - %(config)s - ) RETURNING id; - """) - params = run.model_dump_postgres() - c.execute(query, params) - run_id = c.fetchone()["id"] - run.id = run_id - - def create_tool_run(self, run: NmapRun | RDNSRun | MTRRun): - if type(run) is NmapRun: - return self.create_nmap_run(run) - elif type(run) is RDNSRun: - return self.create_rdns_run(run) - elif type(run) is MTRRun: - return self.create_mtr_run(run) - else: - raise ValueError("unrecognized run type") - - def get_latest_runs_by_tool(self, ip: str) -> dict[ToolName, ToolRun]: - query = """ - SELECT DISTINCT ON (tool_name) * - FROM network_toolrun - WHERE ip = %(ip)s - ORDER BY tool_name, started_at DESC; - """ - params = {"ip": ip} - res = self.pg_config.execute_sql_query(query, params=params) - runs = [ToolRun.model_validate(x) for x in res] - return {r.tool_name: r for r in runs} - - def create_nmap_run(self, run: NmapRun) -> NmapRun: - """ - Insert a PortScan + PortScanPorts from a Pydantic NmapResult. - """ - with self.pg_config.make_connection() as conn, conn.cursor() as c: - self._create_tool_run(run, c) - self.nmap_manager._create(run, c=c) - return run - - def get_nmap_run(self, id: int) -> NmapRun: - query = """ - SELECT tr.*, np.parsed - FROM network_toolrun tr - JOIN network_portscan np ON tr.id = np.run_id - WHERE id = %(id)s - """ - params = {"id": id} - res = self.pg_config.execute_sql_query(query, params)[0] - return NmapRun.model_validate(res) - - def create_rdns_run(self, run: RDNSRun) -> RDNSRun: - """ - Insert a RDnsRun + RDNSResult - """ - with self.pg_config.make_connection() as conn, conn.cursor() as c: - self._create_tool_run(run, c) - self.rdns_manager._create(run, c=c) - return run - - def get_rdns_run(self, id: int) -> RDNSRun: - query = """ - SELECT tr.*, hostnames - FROM network_toolrun tr - JOIN network_rdnsresult np ON tr.id = np.run_id - WHERE id = %(id)s - """ - params = {"id": id} - res = self.pg_config.execute_sql_query(query, params)[0] - parsed = RDNSResult.model_validate( - {"ip": res["ip"], "hostnames": res["hostnames"]} - ) - res["parsed"] = parsed - return RDNSRun.model_validate(res) - - def create_mtr_run(self, run: MTRRun) -> MTRRun: - with self.pg_config.make_connection() as conn, conn.cursor() as c: - self._create_tool_run(run, c) - self.mtr_manager._create(run, c=c) - return run - - def get_mtr_run(self, id: int) -> MTRRun: - query = """ - SELECT tr.*, mtr.parsed, mtr.source_ip, mtr.facility_id - FROM network_toolrun tr - JOIN network_mtr mtr ON tr.id = mtr.run_id - WHERE id = %(id)s - """ - params = {"id": id} - res = self.pg_config.execute_sql_query(query, params)[0] - return MTRRun.model_validate(res) diff --git a/generalresearch/models/cint/question.py b/generalresearch/models/cint/question.py index 0f5453b..89c7871 100644 --- a/generalresearch/models/cint/question.py +++ b/generalresearch/models/cint/question.py @@ -3,11 +3,12 @@ from __future__ import annotations import json from datetime import UTC, datetime from enum import StrEnum -from typing import TYPE_CHECKING, Any, Literal, Self +from typing import Any, Literal, Self from uuid import UUID from pydantic import BaseModel, Field, field_validator, model_validator +from generalresearch.models.cint import CintQuestionIdType from generalresearch.models.custom_types import AwareDatetimeISO from generalresearch.models.definitions import Source from generalresearch.models.string_utils import remove_nbsp @@ -15,12 +16,9 @@ from generalresearch.models.thl.profiling.marketplace import ( MarketplaceQuestion, MarketplaceUserQuestionAnswer, ) - -if TYPE_CHECKING: - from generalresearch.models.cint import CintQuestionIdType - from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestion, - ) +from generalresearch.models.thl.profiling.upk_question import ( + UpkQuestion, +) class CintQuestionType(StrEnum): diff --git a/generalresearch/models/cint/survey.py b/generalresearch/models/cint/survey.py index cd429dd..b2a8935 100644 --- a/generalresearch/models/cint/survey.py +++ b/generalresearch/models/cint/survey.py @@ -4,7 +4,7 @@ import json import logging from datetime import UTC, datetime from decimal import Decimal -from typing import TYPE_CHECKING, Annotated, Any, Literal, Self +from typing import Annotated, Any, Literal, Self from more_itertools import flatten from pydantic import ( @@ -18,6 +18,7 @@ from pydantic import ( ) from generalresearch.locales import Localelator +from generalresearch.models.cint import CintQuestionIdType from generalresearch.models.custom_types import ( AlphaNumStr, AwareDatetimeISO, @@ -31,10 +32,6 @@ from generalresearch.models.thl.survey.condition import ( MarketplaceCondition, ) -if TYPE_CHECKING: - from generalresearch.models.cint import CintQuestionIdType - - logging.basicConfig() logger = logging.getLogger() logger.setLevel(logging.INFO) @@ -76,9 +73,9 @@ class CintQuota(BaseModel): @model_validator(mode="after") def validate_condition_len(self) -> Self: if self.quota_type == "total": - assert ( - self.condition_hashes is None - ), "total quota should not have conditions" + assert self.condition_hashes is None, ( + "total quota should not have conditions" + ) elif self.quota_type == "client": assert len(self.condition_hashes) > 0, "quota must have conditions" return self diff --git a/generalresearch/models/cint/task_collection.py b/generalresearch/models/cint/task_collection.py index 31a0173..4ae8de4 100644 --- a/generalresearch/models/cint/task_collection.py +++ b/generalresearch/models/cint/task_collection.py @@ -1,19 +1,15 @@ from __future__ import annotations -from typing import TYPE_CHECKING - import pandas as pd from pandera.pandas import Check, Column, DataFrameSchema, Index from generalresearch.locales import Localelator +from generalresearch.models.cint.survey import CintSurvey from generalresearch.models.thl.survey.task_collection import ( TaskCollection, create_empty_df_from_schema, ) -if TYPE_CHECKING: - from generalresearch.models.cint.survey import CintSurvey - COUNTRY_ISOS: set[str] = Localelator().get_all_countries() LANGUAGE_ISOS: set[str] = Localelator().get_all_languages() diff --git a/generalresearch/models/dynata/survey.py b/generalresearch/models/dynata/survey.py index 491157b..6ff397f 100644 --- a/generalresearch/models/dynata/survey.py +++ b/generalresearch/models/dynata/survey.py @@ -5,7 +5,7 @@ import logging from datetime import UTC from decimal import Decimal from functools import cached_property -from typing import TYPE_CHECKING, Any, Literal, Self +from typing import Any, Literal, Self from more_itertools import flatten from pydantic import ( @@ -26,7 +26,7 @@ from generalresearch.models.custom_types import ( CoercedStr, DeviceTypes, ) -from generalresearch.models.definitions import Source +from generalresearch.models.definitions import Source, TaskCalculationType from generalresearch.models.dynata import DynataStatus from generalresearch.models.thl.demographics import ( Gender, @@ -37,10 +37,6 @@ from generalresearch.models.thl.survey.condition import ( MarketplaceCondition, ) -if TYPE_CHECKING: - - from generalresearch.models.definitions import TaskCalculationType - logging.basicConfig() logger = logging.getLogger() logger.setLevel(logging.INFO) @@ -136,7 +132,7 @@ class DynataCondition(MarketplaceCondition): if cell["kind"] == "RANGE": d["values"] = [ - f"{cell["range"]["from"] or "inf"}-{cell["range"]["to"] or "inf"}" + f"{cell['range']['from'] or 'inf'}-{cell['range']['to'] or 'inf'}" ] d["value_type"] = ConditionValueType.RANGE return cls.model_validate(d) diff --git a/generalresearch/models/dynata/task_collection.py b/generalresearch/models/dynata/task_collection.py index c6cdc19..e6f0548 100644 --- a/generalresearch/models/dynata/task_collection.py +++ b/generalresearch/models/dynata/task_collection.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Any +from typing import Any import pandas as pd from pandera.pandas import Check, Column, DataFrameSchema, Index @@ -8,14 +8,12 @@ from pandera.pandas import Check, Column, DataFrameSchema, Index from generalresearch.locales import Localelator from generalresearch.models.definitions import TaskCalculationType from generalresearch.models.dynata import DynataStatus +from generalresearch.models.dynata.survey import DynataSurvey from generalresearch.models.thl.survey.task_collection import ( TaskCollection, create_empty_df_from_schema, ) -if TYPE_CHECKING: - from generalresearch.models.dynata.survey import DynataSurvey - COUNTRY_ISOS = Localelator().get_all_countries() LANGUAGE_ISOS = Localelator().get_all_languages() diff --git a/generalresearch/models/innovate/question.py b/generalresearch/models/innovate/question.py index 6423399..4af0639 100644 --- a/generalresearch/models/innovate/question.py +++ b/generalresearch/models/innovate/question.py @@ -4,21 +4,19 @@ from __future__ import annotations import json import logging from enum import StrEnum -from typing import TYPE_CHECKING, Any, Literal +from typing import Any, Literal from pydantic import BaseModel, Field, ValidationError, field_validator, model_validator from generalresearch.models.definitions import Source +from generalresearch.models.innovate import InnovateQuestionID from generalresearch.models.thl.profiling.marketplace import ( MarketplaceQuestion, MarketplaceUserQuestionAnswer, ) - -if TYPE_CHECKING: - from generalresearch.models.innovate import InnovateQuestionID - from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestion, - ) +from generalresearch.models.thl.profiling.upk_question import ( + UpkQuestion, +) logging.basicConfig() logger = logging.getLogger() diff --git a/generalresearch/models/innovate/survey.py b/generalresearch/models/innovate/survey.py index 7232228..3ca990c 100644 --- a/generalresearch/models/innovate/survey.py +++ b/generalresearch/models/innovate/survey.py @@ -6,7 +6,6 @@ from datetime import UTC, date from decimal import Decimal from functools import cached_property from typing import ( - TYPE_CHECKING, Annotated, Any, Literal, @@ -33,12 +32,14 @@ from generalresearch.models.custom_types import ( from generalresearch.models.definitions import ( LogicalOperator, Source, + TaskCalculationType, ) from generalresearch.models.innovate import ( InnovateDuplicateCheckLevel, InnovateQuotaStatus, InnovateStatus, ) +from generalresearch.models.innovate.question import InnovateQuestionID from generalresearch.models.thl.demographics import Gender from generalresearch.models.thl.survey import MarketplaceTask from generalresearch.models.thl.survey.condition import ( @@ -46,13 +47,6 @@ from generalresearch.models.thl.survey.condition import ( MarketplaceCondition, ) -if TYPE_CHECKING: - - from generalresearch.models.definitions import ( - TaskCalculationType, - ) - from generalresearch.models.innovate.question import InnovateQuestionID - logging.basicConfig() logger = logging.getLogger() logger.setLevel(logging.INFO) diff --git a/generalresearch/models/innovate/task_collection.py b/generalresearch/models/innovate/task_collection.py index a647d30..7bf9d0f 100644 --- a/generalresearch/models/innovate/task_collection.py +++ b/generalresearch/models/innovate/task_collection.py @@ -1,20 +1,16 @@ from __future__ import annotations -from typing import TYPE_CHECKING - import pandas as pd from pandera.pandas import Check, Column, DataFrameSchema, Index from generalresearch.locales import Localelator from generalresearch.models.innovate import InnovateStatus +from generalresearch.models.innovate.survey import InnovateSurvey from generalresearch.models.thl.survey.task_collection import ( TaskCollection, create_empty_df_from_schema, ) -if TYPE_CHECKING: - from generalresearch.models.innovate.survey import InnovateSurvey - COUNTRY_ISOS: set[str] = Localelator().get_all_countries() LANGUAGE_ISOS: set[str] = Localelator().get_all_languages() diff --git a/generalresearch/models/legacy/offerwall.py b/generalresearch/models/legacy/offerwall.py index 013b506..da28663 100644 --- a/generalresearch/models/legacy/offerwall.py +++ b/generalresearch/models/legacy/offerwall.py @@ -1,34 +1,27 @@ from __future__ import annotations -from typing import TYPE_CHECKING - from pydantic import BaseModel, ConfigDict, Field, NonNegativeInt from generalresearch.models.custom_types import UUIDStr +from generalresearch.models.legacy.bucket import ( + BucketBase, + MarketplaceBucket, + OneShotOfferwallBucket, + OneShotSoftPairOfferwallBucket, + SingleEntryBucket, + SoftPairBucket, + TimeBucksBucket, + TopNBucket, + TopNPlusBucket, + TopNPlusRecontactBucket, + WXETOfferwallBucket, +) from generalresearch.models.legacy.definitions import OfferwallReason from generalresearch.models.thl.payout_format import ( PayoutFormatField, + PayoutFormatType, ) - -if TYPE_CHECKING: - - from generalresearch.models.legacy.bucket import ( - BucketBase, - MarketplaceBucket, - OneShotOfferwallBucket, - OneShotSoftPairOfferwallBucket, - SingleEntryBucket, - SoftPairBucket, - TimeBucksBucket, - TopNBucket, - TopNPlusBucket, - TopNPlusRecontactBucket, - WXETOfferwallBucket, - ) - from generalresearch.models.thl.payout_format import ( - PayoutFormatType, - ) - from generalresearch.models.thl.profiling.upk_question import UpkQuestion +from generalresearch.models.thl.profiling.upk_question import UpkQuestion """ Not Done: diff --git a/generalresearch/models/legacy/questions.py b/generalresearch/models/legacy/questions.py index e6803f0..caa6aae 100644 --- a/generalresearch/models/legacy/questions.py +++ b/generalresearch/models/legacy/questions.py @@ -17,17 +17,15 @@ from sentry_sdk import capture_exception from generalresearch.models.custom_types import UUIDStr from generalresearch.models.legacy.api_status import StatusResponse +from generalresearch.models.thl.profiling.upk_question import ( + UpkQuestionOut, +) +from generalresearch.models.thl.session import Wall +from generalresearch.models.thl.user import User if TYPE_CHECKING: - from generalresearch.managers.thl.user_manager.user_manager import ( - UserManager, - ) + from generalresearch.managers.thl.user_manager.user_manager import UserManager from generalresearch.managers.thl.wall import WallManager - from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestionOut, - ) - from generalresearch.models.thl.session import Wall - from generalresearch.models.thl.user import User class UpkQuestionResponse(StatusResponse): diff --git a/generalresearch/models/lucid/question.py b/generalresearch/models/lucid/question.py index c1b9e52..288f0d2 100644 --- a/generalresearch/models/lucid/question.py +++ b/generalresearch/models/lucid/question.py @@ -2,20 +2,18 @@ from __future__ import annotations import logging from enum import StrEnum -from typing import TYPE_CHECKING, Any, Literal, Self +from typing import Any, Literal, Self from pydantic import BaseModel, Field, field_validator, model_validator from generalresearch.models.definitions import Source +from generalresearch.models.lucid import LucidQuestionIdType from generalresearch.models.thl.profiling.marketplace import ( MarketplaceQuestion, ) - -if TYPE_CHECKING: - from generalresearch.models.lucid import LucidQuestionIdType - from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestion, - ) +from generalresearch.models.thl.profiling.upk_question import ( + UpkQuestion, +) logging.basicConfig() logger = logging.getLogger() diff --git a/generalresearch/models/lucid/survey.py b/generalresearch/models/lucid/survey.py index 02b31ab..bca471b 100644 --- a/generalresearch/models/lucid/survey.py +++ b/generalresearch/models/lucid/survey.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Any, Self +from typing import Any, Self from pydantic import BaseModel, ConfigDict, Field, NonNegativeInt @@ -11,14 +11,12 @@ from generalresearch.models.custom_types import ( UUIDStr, ) from generalresearch.models.definitions import Source +from generalresearch.models.thl.locales import CountryISO, LanguageISO from generalresearch.models.thl.survey.condition import ( ConditionValueType, MarketplaceCondition, ) -if TYPE_CHECKING: - from generalresearch.models.thl.locales import CountryISO, LanguageISO - class LucidCondition(MarketplaceCondition): model_config = ConfigDict(populate_by_name=True, frozen=False, extra="ignore") diff --git a/generalresearch/models/morning/question.py b/generalresearch/models/morning/question.py index 909992f..7c676fb 100644 --- a/generalresearch/models/morning/question.py +++ b/generalresearch/models/morning/question.py @@ -1,20 +1,18 @@ import json from enum import StrEnum -from typing import TYPE_CHECKING, Any, Literal, Self +from typing import Any, Literal, Self from uuid import UUID from pydantic import BaseModel, Field, field_validator, model_validator from generalresearch.locales import Localelator from generalresearch.models.definitions import Source +from generalresearch.models.morning import MorningQuestionID from generalresearch.models.thl.profiling.marketplace import ( MarketplaceQuestion, MarketplaceUserQuestionAnswer, ) -if TYPE_CHECKING: - from generalresearch.models.morning import MorningQuestionID - # todo: we could validate that the country_iso / language_iso exists ... locale_helper = Localelator() diff --git a/generalresearch/models/morning/survey.py b/generalresearch/models/morning/survey.py index 01cc4ff..8738631 100644 --- a/generalresearch/models/morning/survey.py +++ b/generalresearch/models/morning/survey.py @@ -6,7 +6,6 @@ from datetime import UTC from decimal import Decimal from functools import cached_property from typing import ( - TYPE_CHECKING, Annotated, Any, Literal, @@ -30,24 +29,20 @@ from generalresearch.models.custom_types import ( UUIDStrCoerce, ) from generalresearch.models.definitions import Source -from generalresearch.models.morning import MorningStatus +from generalresearch.models.morning import MorningQuestionID, MorningStatus +from generalresearch.models.morning.question import MorningQuestion from generalresearch.models.thl.demographics import Gender +from generalresearch.models.thl.locales import ( + CountryISO, + CountryISOs, + LanguageISOs, +) from generalresearch.models.thl.survey import MarketplaceTask from generalresearch.models.thl.survey.condition import ( ConditionValueType, MarketplaceCondition, ) -if TYPE_CHECKING: - - from generalresearch.models.morning import MorningQuestionID - from generalresearch.models.morning.question import MorningQuestion - from generalresearch.models.thl.locales import ( - CountryISO, - CountryISOs, - LanguageISOs, - ) - logging.basicConfig() logger = logging.getLogger() logger.setLevel(logging.INFO) diff --git a/generalresearch/models/morning/task_collection.py b/generalresearch/models/morning/task_collection.py index 1117937..9303a2f 100644 --- a/generalresearch/models/morning/task_collection.py +++ b/generalresearch/models/morning/task_collection.py @@ -1,20 +1,16 @@ from __future__ import annotations -from typing import TYPE_CHECKING - import pandas as pd from pandera.pandas import Check, Column, DataFrameSchema, Index from generalresearch.locales import Localelator from generalresearch.models.morning import MorningStatus +from generalresearch.models.morning.survey import MorningBid from generalresearch.models.thl.survey.task_collection import ( TaskCollection, create_empty_df_from_schema, ) -if TYPE_CHECKING: - from generalresearch.models.morning.survey import MorningBid - COUNTRY_ISOS: set[str] = Localelator().get_all_countries() LANGUAGE_ISOS: set[str] = Localelator().get_all_languages() diff --git a/generalresearch/models/network/__init__.py b/generalresearch/models/network/__init__.py deleted file mode 100644 index e69de29..0000000 diff --git a/generalresearch/models/network/definitions.py b/generalresearch/models/network/definitions.py deleted file mode 100644 index 2e1ab91..0000000 --- a/generalresearch/models/network/definitions.py +++ /dev/null @@ -1,70 +0,0 @@ -from __future__ import annotations - -from enum import StrEnum -from ipaddress import ip_address, ip_network - -CGNAT_NET = ip_network("100.64.0.0/10") - - -class IPProtocol(StrEnum): - TCP = "tcp" - UDP = "udp" - SCTP = "sctp" - IP = "ip" - ICMP = "icmp" - ICMPv6 = "icmpv6" - - def to_number(self) -> int: - # https://www.iana.org/assignments/protocol-numbers/protocol-numbers.xhtml - return { - self.TCP: 6, - self.UDP: 17, - self.SCTP: 132, - self.IP: 4, - self.ICMP: 1, - self.ICMPv6: 58, - }[self] - - -class IPKind(StrEnum): - PUBLIC = "public" - PRIVATE = "private" - CGNAT = "carrier_nat" - LOOPBACK = "loopback" - LINK_LOCAL = "link_local" - MULTICAST = "multicast" - RESERVED = "reserved" - UNSPECIFIED = "unspecified" - - -def get_ip_kind(ip: str | None) -> IPKind | None: - if not ip: - return None - - ip_obj = ip_address(ip) - - if ip_obj in CGNAT_NET: - return IPKind.CGNAT - - if ip_obj.is_loopback: - return IPKind.LOOPBACK - - if ip_obj.is_link_local: - return IPKind.LINK_LOCAL - - if ip_obj.is_multicast: - return IPKind.MULTICAST - - if ip_obj.is_unspecified: - return IPKind.UNSPECIFIED - - if ip_obj.is_private: - return IPKind.PRIVATE - - if ip_obj.is_reserved: - return IPKind.RESERVED - - if ip_obj.is_global: - return IPKind.PUBLIC - - return None diff --git a/generalresearch/models/network/label.py b/generalresearch/models/network/label.py deleted file mode 100644 index c27f36f..0000000 --- a/generalresearch/models/network/label.py +++ /dev/null @@ -1,129 +0,0 @@ -from __future__ import annotations - -import ipaddress -from enum import StrEnum -from ipaddress import IPv4Network, IPv6Network -from typing import TYPE_CHECKING - -from pydantic import ( - BaseModel, - ConfigDict, - Field, - IPvAnyNetwork, - computed_field, - field_validator, -) - -from generalresearch.models.custom_types import now_utc_factory - -if TYPE_CHECKING: - from generalresearch.models.custom_types import AwareDatetimeISO - - -class IPTrustClass(StrEnum): - TRUSTED = "trusted" - UNTRUSTED = "untrusted" - # Note: use case of unknown is for e.g. Spur says this IP is a residential proxy - # on 2026-1-1, and then has no annotation a month later. It doesn't mean - # the IP is TRUSTED, but we want to record that Spur now doesn't claim UNTRUSTED. - UNKNOWN = "unknown" - - -class IPLabelKind(StrEnum): - # --- UNTRUSTED --- - RESIDENTIAL_PROXY = "residential_proxy" - DATACENTER_PROXY = "datacenter_proxy" - ISP_PROXY = "isp_proxy" - MOBILE_PROXY = "mobile_proxy" - PROXY = "proxy" - HOSTING = "hosting" - VPN = "vpn" - RELAY = "relay" - TOR_EXIT = "tor_exit" - BAD_ACTOR = "bad_actor" - # --- TRUSTED --- - TRUSTED_USER = "trusted_user" - # --- UNKNOWN --- - UNKNOWN = "unknown" - - -class IPLabelSource(StrEnum): - # We got this IP from our own use of a proxy service - INTERNAL_USE = "internal_use" - - # An external "security" service flagged this IP - GRIP = "grip" - SPUR = "spur" - IPINFO = "ipinfo" - MAXMIND = "maxmind" - - MANUAL = "manual" - - -class IPLabel(BaseModel): - """ - Stores *ground truth* about an IP at a specific time. - To be used for model training and evaluation. - """ - - model_config = ConfigDict(validate_assignment=True) - - ip: IPvAnyNetwork = Field() - - labeled_at: AwareDatetimeISO = Field(default_factory=now_utc_factory) - created_at: AwareDatetimeISO | None = Field(default=None) - - label_kind: IPLabelKind = Field() - source: IPLabelSource = Field() - - confidence: float = Field(default=1.0, ge=0.0, le=1.0) - - # Optionally, if this is untrusted, which service is providing the proxy/vpn service - provider: str | None = Field( - default=None, examples=["geonode", "gecko"], max_length=128 - ) - - metadata: IPLabelMetadata | None = Field(default=None) - - @field_validator("ip", mode="before") - @classmethod - def normalize_and_validate_network( - cls, v: IPvAnyNetwork - ) -> IPv4Network | IPv6Network | None: - net = ipaddress.ip_network(address=v, strict=False) - - if isinstance(net, ipaddress.IPv6Network) and net.prefixlen > 64: - raise ValueError("IPv6 network must be /64 or larger") - - return net - - @field_validator("provider", mode="before") - @classmethod - def provider_format(cls, v: str | None) -> str | None: - if v is None: - return v - return v.lower().strip() - - @computed_field() - @property - def trust_class(self) -> IPTrustClass: - if self.label_kind == IPLabelKind.UNKNOWN: - return IPTrustClass.UNKNOWN - if self.label_kind == IPLabelKind.TRUSTED_USER: - return IPTrustClass.TRUSTED - return IPTrustClass.UNTRUSTED - - def model_dump_postgres(self): - d = self.model_dump(mode="json") - d["metadata"] = self.metadata.model_dump_json() if self.metadata else None - return d - - -class IPLabelMetadata(BaseModel): - """ - To be expanded. Just for storing some things from Spur for now - """ - - model_config = ConfigDict(validate_assignment=True, extra="allow") - - services: list[str] | None = Field(min_length=1, examples=[["RDP"]]) diff --git a/generalresearch/models/network/mtr/__init__.py b/generalresearch/models/network/mtr/__init__.py deleted file mode 100644 index e69de29..0000000 diff --git a/generalresearch/models/network/mtr/command.py b/generalresearch/models/network/mtr/command.py deleted file mode 100644 index 7e74f20..0000000 --- a/generalresearch/models/network/mtr/command.py +++ /dev/null @@ -1,76 +0,0 @@ -from __future__ import annotations - -import subprocess -from typing import TYPE_CHECKING - -from generalresearch.models.network.definitions import IPProtocol -from generalresearch.models.network.mtr.parser import parse_mtr_output - -if TYPE_CHECKING: - from generalresearch.models.network.mtr.result import MTRResult - from generalresearch.models.network.tool_run_command import MTRRunCommand - -SUPPORTED_PROTOCOLS = { - IPProtocol.TCP, - IPProtocol.UDP, - IPProtocol.SCTP, - IPProtocol.ICMP, -} -PROTOCOLS_W_PORT = {IPProtocol.TCP, IPProtocol.UDP, IPProtocol.SCTP} - - -def build_mtr_command( - ip: str, - protocol: IPProtocol | None = None, - port: int | None = None, - report_cycles: int | None = 10, -) -> str: - # https://manpages.ubuntu.com/manpages/focal/man8/mtr.8.html - # e.g. "mtr -r -c 2 -b -z -j -T -P 443 74.139.70.149" - args = ["mtr", "--report", "--show-ips", "--aslookup", "--json"] - if report_cycles is not None: - args.extend(["-c", str(int(report_cycles))]) - if port is not None: - if protocol is None: - protocol = IPProtocol.TCP - assert protocol in PROTOCOLS_W_PORT, "port only allowed for TCP/SCTP/UDP traces" - args.extend(["--port", str(int(port))]) - if protocol: - assert protocol in SUPPORTED_PROTOCOLS, f"unsupported protocol: {protocol}" - # default is ICMP (no args) - arg_map = { - IPProtocol.TCP: "--tcp", - IPProtocol.UDP: "--udp", - IPProtocol.SCTP: "--sctp", - } - if protocol in arg_map: - args.append(arg_map[protocol]) - args.append(ip) - return " ".join(args) - - -def get_mtr_version() -> str: - proc = subprocess.run( - ["mtr", "-v"], - capture_output=True, - text=True, - check=False, - ) - # e.g. mtr 0.95 - ver_str = proc.stdout.strip() - return ver_str.split(" ", 1)[1] - - -def run_mtr(config: MTRRunCommand) -> MTRResult: - cmd = config.to_command_str() - args = cmd.split(" ") - proc = subprocess.run( - args, - capture_output=True, - text=True, - check=False, - ) - raw = proc.stdout.strip() - return parse_mtr_output( - raw, protocol=config.options.protocol, port=config.options.port - ) diff --git a/generalresearch/models/network/mtr/execute.py b/generalresearch/models/network/mtr/execute.py deleted file mode 100644 index 5ab7632..0000000 --- a/generalresearch/models/network/mtr/execute.py +++ /dev/null @@ -1,60 +0,0 @@ -from __future__ import annotations - -from datetime import UTC, datetime -from uuid import uuid4 - -from generalresearch.models.custom_types import UUIDStr -from generalresearch.models.network.definitions import IPProtocol -from generalresearch.models.network.mtr.command import ( - get_mtr_version, - run_mtr, -) -from generalresearch.models.network.tool_run import ( - MTRRun, - Status, - ToolClass, - ToolName, -) -from generalresearch.models.network.tool_run_command import ( - MTRRunCommand, - MTRRunCommandOptions, -) -from generalresearch.models.network.utils import get_source_ip - - -def execute_mtr( - ip: str, - scan_group_id: UUIDStr | None = None, - protocol: IPProtocol | None = IPProtocol.ICMP, - port: int | None = None, - report_cycles: int = 10, -) -> MTRRun: - config = MTRRunCommand( - options=MTRRunCommandOptions( - ip=ip, - report_cycles=report_cycles, - protocol=protocol, - port=port, - ), - ) - - started_at = datetime.now(tz=UTC) - tool_version = get_mtr_version() - result = run_mtr(config) - finished_at = datetime.now(tz=UTC) - - return MTRRun( - tool_name=ToolName.MTR, - tool_class=ToolClass.TRACEROUTE, - tool_version=tool_version, - status=Status.SUCCESS, - ip=ip, - started_at=started_at, - finished_at=finished_at, - raw_command=config.to_command_str(), - scan_group_id=scan_group_id or uuid4().hex, - config=config, - parsed=result, - source_ip=get_source_ip(), - facility_id=1, - ) diff --git a/generalresearch/models/network/mtr/parser.py b/generalresearch/models/network/mtr/parser.py deleted file mode 100644 index 30c22bf..0000000 --- a/generalresearch/models/network/mtr/parser.py +++ /dev/null @@ -1,19 +0,0 @@ -import json -from typing import Any - -from generalresearch.models.network.definitions import IPProtocol -from generalresearch.models.network.mtr.result import MTRResult - - -def parse_mtr_output(raw: str, port: int, protocol: IPProtocol) -> MTRResult: - data = parse_mtr_raw_output(raw) - data["port"] = port - data["protocol"] = protocol - return MTRResult.model_validate(data) - - -def parse_mtr_raw_output(raw: str) -> dict[str, Any]: - data = json.loads(raw)["report"] - data.update(data.pop("mtr")) - data["hops"] = data.pop("hubs") - return data diff --git a/generalresearch/models/network/mtr/result.py b/generalresearch/models/network/mtr/result.py deleted file mode 100644 index d17136c..0000000 --- a/generalresearch/models/network/mtr/result.py +++ /dev/null @@ -1,175 +0,0 @@ -from __future__ import annotations - -import re -from functools import cached_property -from ipaddress import ip_address -from typing import TYPE_CHECKING - -import tldextract -from pydantic import ( - BaseModel, - ConfigDict, - Field, - computed_field, - field_validator, - model_validator, -) - -from generalresearch.models.network.definitions import ( - IPProtocol, - get_ip_kind, -) - -if TYPE_CHECKING: - from generalresearch.models.network.definitions import IPKind - -HOST_RE = re.compile(r"^(?P.+?) \((?P[^)]+)\)$") - - -class MTRHop(BaseModel): - model_config = ConfigDict(populate_by_name=True) - - hop: int = Field(alias="count") - host: str - asn: int | None = Field(default=None, alias="ASN") - - loss_pct: float = Field(alias="Loss%") - sent: int = Field(alias="Snt") - - last_ms: float = Field(alias="Last") - avg_ms: float = Field(alias="Avg") - best_ms: float = Field(alias="Best") - worst_ms: float = Field(alias="Wrst") - stdev_ms: float = Field(alias="StDev") - - hostname: str | None = Field( - default=None, examples=["fixed-187-191-8-145.totalplay.net"] - ) - ip: str | None = None - - @field_validator("asn", mode="before") - @classmethod - def normalize_asn(cls, v: str): - if v is None or v == "AS???": - return None - if type(v) is int: - return v - return int(v.replace("AS", "")) - - @model_validator(mode="after") - def parse_host(self): - host = self.host.strip() - - # hostname (ip) - m = HOST_RE.match(host) - if m: - self.hostname = m.group("hostname") - self.ip = m.group("ip") - return self - - # ip only - try: - ip_address(host) - self.ip = host - self.hostname = None - return self - except ValueError: - pass - - # hostname only - self.hostname = host - self.ip = None - return self - - @cached_property - def ip_kind(self) -> IPKind | None: - return get_ip_kind(self.ip) - - @cached_property - def icmp_rate_limited(self): - if self.avg_ms == 0: - return False - return self.stdev_ms > self.avg_ms or self.worst_ms > self.best_ms * 10 - - @computed_field(examples=["totalplay.net"]) - @cached_property - def domain(self) -> str | None: - if self.hostname: - return tldextract.extract(self.hostname).top_domain_under_public_suffix - - def model_dump_postgres(self, run_id: int): - # Writes for the network_mtrhop table - d = {"mtr_run_id": run_id} - data = self.model_dump( - mode="json", - include={ - "hop", - "ip", - "domain", - "asn", - }, - ) - d.update(data) - return d - - -class MTRResult(BaseModel): - model_config = ConfigDict(populate_by_name=True) - - source: str = Field(description="Hostname of the system running mtr.", alias="src") - destination: str = Field( - description="Destination hostname or IP being traced.", alias="dst" - ) - tos: int = Field(description="IP Type-of-Service (TOS) value used for probes.") - tests: int = Field(description="Number of probes sent per hop.") - psize: int = Field(description="Probe packet size in bytes.") - bitpattern: str = Field(description="Payload byte pattern used in probes (hex).") - - # Protocol used for the traceroute - protocol: IPProtocol = Field(default=IPProtocol.ICMP) - # The target port number for TCP/SCTP/UDP traces - port: int | None = Field(default=None) - - hops: list[MTRHop] = Field() - - def model_dump_postgres(self): - # Writes for the network_mtr table - d = self.model_dump( - mode="json", - include={"port"}, - ) - d["protocol"] = self.protocol.to_number() - d["parsed"] = self.model_dump_json(indent=0) - return d - - def print_report(self) -> None: - print( - f"MTR Report → {self.destination} {self.protocol.name} {self.port or ''}\n" - ) - host_max_len = max(len(h.host) for h in self.hops) - - header = ( - f"{'Hop':>3} " - f"{'Host':<{host_max_len}} " - f"{'Kind':<10} " - f"{'ASN':<8} " - f"{'Loss%':>6} {'Sent':>5} " - f"{'Last':>7} {'Avg':>7} {'Best':>7} {'Worst':>7} {'StDev':>7}" - ) - print(header) - print("-" * len(header)) - - for hop in self.hops: - print( - f"{hop.hop:>3} " - f"{hop.host:<{host_max_len}} " - f"{hop.ip_kind or '???':<10} " - f"{hop.asn or '???':<8} " - f"{hop.loss_pct:6.1f} " - f"{hop.sent:5d} " - f"{hop.last_ms:7.1f} " - f"{hop.avg_ms:7.1f} " - f"{hop.best_ms:7.1f} " - f"{hop.worst_ms:7.1f} " - f"{hop.stdev_ms:7.1f}" - ) diff --git a/generalresearch/models/network/nmap/__init__.py b/generalresearch/models/network/nmap/__init__.py deleted file mode 100644 index e69de29..0000000 diff --git a/generalresearch/models/network/nmap/command.py b/generalresearch/models/network/nmap/command.py deleted file mode 100644 index 3509b8d..0000000 --- a/generalresearch/models/network/nmap/command.py +++ /dev/null @@ -1,51 +0,0 @@ -from __future__ import annotations - -import subprocess -from typing import TYPE_CHECKING - -from generalresearch.models.network.nmap.parser import parse_nmap_xml - -if TYPE_CHECKING: - from generalresearch.models.network.nmap.result import NmapResult - from generalresearch.models.network.tool_run_command import NmapRunCommand - - -def build_nmap_command( - ip: str, - no_ping: bool = True, - enable_advanced: bool = True, - timing: int = 4, - ports: str | None = None, - top_ports: int | None = None, -) -> str: - # e.g. "nmap -Pn -T4 -A --top-ports 1000 -oX - scanme.nmap.org" - # https://linux.die.net/man/1/nmap - args = ["nmap"] - assert 0 <= timing <= 5 - args.append(f"-T{timing}") - if no_ping: - args.append("-Pn") - if enable_advanced: - args.append("-A") - if ports is not None: - assert top_ports is None - args.extend(["-p", ports]) - if top_ports is not None: - assert ports is None - args.extend(["--top-ports", str(top_ports)]) - - args.extend(["-oX", "-", ip]) - return " ".join(args) - - -def run_nmap(config: NmapRunCommand) -> NmapResult: - cmd = config.to_command_str() - args = cmd.split(" ") - proc = subprocess.run( - args, - capture_output=True, - text=True, - check=False, - ) - raw = proc.stdout.strip() - return parse_nmap_xml(raw) diff --git a/generalresearch/models/network/nmap/execute.py b/generalresearch/models/network/nmap/execute.py deleted file mode 100644 index 09ec28b..0000000 --- a/generalresearch/models/network/nmap/execute.py +++ /dev/null @@ -1,56 +0,0 @@ -from __future__ import annotations - -from uuid import uuid4 - -from generalresearch.models.custom_types import UUIDStr -from generalresearch.models.network.nmap.command import run_nmap -from generalresearch.models.network.tool_run import ( - NmapRun, - Status, - ToolClass, - ToolName, -) -from generalresearch.models.network.tool_run_command import ( - NmapRunCommand, - NmapRunCommandOptions, -) - - -def execute_nmap( - ip: str, - top_ports: int | None = 1000, - ports: str | None = None, - no_ping: bool = True, - enable_advanced: bool = True, - timing: int = 4, - scan_group_id: UUIDStr | None = None, -) -> NmapRun: - config = NmapRunCommand( - options=NmapRunCommandOptions( - top_ports=top_ports, - ports=ports, - no_ping=no_ping, - enable_advanced=enable_advanced, - timing=timing, - ip=ip, - ) - ) - result = run_nmap(config) - assert result.exit_status == "success" - assert result.target_ip == ip, f"{result.target_ip=}, {ip=}" - assert result.command_line == config.to_command_str() - - run = NmapRun( - tool_name=ToolName.NMAP, - tool_class=ToolClass.PORT_SCAN, - tool_version=result.version, - status=Status.SUCCESS, - ip=ip, - started_at=result.started_at, - finished_at=result.finished_at, - raw_command=result.command_line, - scan_group_id=scan_group_id or uuid4().hex, - config=config, - parsed=result, - ) - return run diff --git a/generalresearch/models/network/nmap/parser.py b/generalresearch/models/network/nmap/parser.py deleted file mode 100644 index 866b4bd..0000000 --- a/generalresearch/models/network/nmap/parser.py +++ /dev/null @@ -1,414 +0,0 @@ -from __future__ import annotations - -import xml.etree.ElementTree as ET -from datetime import UTC, datetime -from typing import Any - -from generalresearch.models.network.definitions import IPProtocol -from generalresearch.models.network.nmap.result import ( - NmapHostname, - NmapHostScript, - NmapHostState, - NmapHostStatusReason, - NmapOSClass, - NmapOSMatch, - NmapPort, - NmapPortStats, - NmapResult, - NmapScanInfo, - NmapScanType, - NmapScript, - NmapService, - NmapTrace, - NmapTraceHop, - PortState, - PortStateReason, -) - - -class NmapParserException(Exception): - def __init__(self, msg): - self.msg = msg - - def __str__(self): - return self.msg - - -class NmapXmlParser: - """ - Example: https://nmap.org/book/output-formats-xml-output.html - Full DTD: https://nmap.org/book/nmap-dtd.html - """ - - @classmethod - def parse_xml(cls, nmap_data: str) -> NmapResult: - """ - Expects a full nmap scan report. - """ - - try: - root = ET.fromstring(nmap_data) - except ET.ParseError as e: - emsg = f"Wrong XML structure: cannot parse data: {e}" - raise NmapParserException(emsg) - - if root.tag != "nmaprun": - raise NmapParserException("Unpexpected data structure for XML " "root node") - return cls._parse_xml_nmaprun(root) - - @classmethod - def _parse_xml_nmaprun(cls, root: ET.Element) -> NmapResult: - """ - This method parses out a full nmap scan report from its XML root - node: . We expect there is only 1 host in this report! - - :param root: Element from xml.ElementTree (top of XML the document) - """ - cls._validate_nmap_root(root) - host_count = len(root.findall(".//host")) - assert host_count == 1, f"Expected 1 host, got {host_count}" - - xml_str = ET.tostring(root, encoding="unicode").replace("\n", "") - nmap_data = {"raw_xml": xml_str} - nmap_data.update(cls._parse_nmaprun(root)) - - nmap_data["scan_infos"] = [ - cls._parse_scaninfo(scaninfo_el) - for scaninfo_el in root.findall(".//scaninfo") - ] - - nmap_data.update(cls._parse_runstats(root)) - - nmap_data.update(cls._parse_xml_host(root.find(".//host"))) - - return NmapResult.model_validate(nmap_data) - - @classmethod - def _validate_nmap_root(cls, root: ET.Element) -> None: - allowed = { - "scaninfo", - "host", - "runstats", - "verbose", - "debugging", - "taskprogress", - } - - found = {child.tag for child in root} - unexpected = found - allowed - if unexpected: - raise ValueError( - f"Unexpected top-level tags in nmap XML: {sorted(unexpected)}" - ) - - @classmethod - def _parse_scaninfo(cls, scaninfo_el: ET.Element) -> NmapScanInfo: - data = {} - data["type"] = NmapScanType(scaninfo_el.attrib["type"]) - data["protocol"] = IPProtocol(scaninfo_el.attrib["protocol"]) - data["num_services"] = scaninfo_el.attrib["numservices"] - data["services"] = scaninfo_el.attrib["services"] - return NmapScanInfo.model_validate(data) - - @classmethod - def _parse_runstats(cls, root: ET.Element) -> dict: - runstats = root.find("runstats") - if runstats is None: - return {} - - finished = runstats.find("finished") - if finished is None: - return {} - - finished_at = None - ts = finished.attrib.get("time") - if ts: - finished_at = datetime.fromtimestamp(int(ts), tz=UTC) - - return { - "finished_at": finished_at, - "exit_status": finished.attrib.get("exit"), - } - - @classmethod - def _parse_nmaprun(cls, nmaprun_el: ET.Element) -> dict: - nmap_data = {} - nmaprun = dict(nmaprun_el.attrib) - nmap_data["command_line"] = nmaprun["args"] - nmap_data["started_at"] = datetime.fromtimestamp( - float(nmaprun["start"]), tz=UTC - ) - nmap_data["version"] = nmaprun["version"] - nmap_data["xmloutputversion"] = nmaprun["xmloutputversion"] - return nmap_data - - @classmethod - def _parse_xml_host(cls, host_el: ET.Element) -> dict: - """ - Receives a XML tag representing a scanned host with - its services. - """ - data = {} - - # - status_el = host_el.find("status") - data["host_state"] = NmapHostState(status_el.attrib["state"]) - data["host_state_reason"] = NmapHostStatusReason(status_el.attrib["reason"]) - host_state_reason_ttl = status_el.attrib.get("reason_ttl") - if host_state_reason_ttl: - data["host_state_reason_ttl"] = int(host_state_reason_ttl) - - #
- address_el = host_el.find("address") - data["target_ip"] = address_el.attrib["addr"] - - data["hostnames"] = cls._parse_hostnames(host_el.find("hostnames")) - - data["ports"], data["port_stats"] = cls._parse_xml_ports(host_el.find("ports")) - - uptime = host_el.find("uptime") - if uptime is not None: - data["uptime_seconds"] = int(uptime.attrib["seconds"]) - - distance = host_el.find("distance") - if distance is not None: - data["distance"] = int(distance.attrib["value"]) - - tcpsequence = host_el.find("tcpsequence") - if tcpsequence is not None: - data["tcp_sequence_index"] = int(tcpsequence.attrib["index"]) - data["tcp_sequence_difficulty"] = tcpsequence.attrib["difficulty"] - ipidsequence = host_el.find("ipidsequence") - if ipidsequence is not None: - data["ipid_sequence_class"] = ipidsequence.attrib["class"] - tcptssequence = host_el.find("tcptssequence") - if tcptssequence is not None: - data["tcp_timestamp_class"] = tcptssequence.attrib["class"] - - times_elem = host_el.find("times") - if times_elem is not None: - data.update( - { - "srtt_us": int(times_elem.attrib.get("srtt", 0)) or None, - "rttvar_us": int(times_elem.attrib.get("rttvar", 0)) or None, - "timeout_us": int(times_elem.attrib.get("to", 0)) or None, - } - ) - - hostscripts_el = host_el.find("hostscript") - if hostscripts_el is not None: - data["host_scripts"] = [ - NmapHostScript(id=el.attrib["id"], output=el.attrib.get("output")) - for el in hostscripts_el.findall("script") - ] - - data["os_matches"] = cls._parse_os_matches(host_el) - - data["trace"] = cls._parse_trace(host_el) - - return data - - @classmethod - def _parse_os_matches(cls, host_el: ET.Element) -> list[NmapOSMatch] | None: - os_elem = host_el.find("os") - if os_elem is None: - return None - - matches: list[NmapOSMatch] = [] - - for m in os_elem.findall("osmatch"): - classes: list[NmapOSClass] = [] - - for c in m.findall("osclass"): - cpes = [e.text.strip() for e in c.findall("cpe") if e.text] - - classes.append( - NmapOSClass( - vendor=c.attrib.get("vendor"), - osfamily=c.attrib.get("osfamily"), - osgen=c.attrib.get("osgen"), - accuracy=( - int(c.attrib["accuracy"]) - if "accuracy" in c.attrib - else None - ), - cpe=cpes or None, - ) - ) - - matches.append( - NmapOSMatch( - name=m.attrib["name"], - accuracy=int(m.attrib["accuracy"]), - classes=classes, - ) - ) - - return matches or None - - @classmethod - def _parse_hostnames(cls, hostnames_el: ET.Element) -> list[NmapHostname]: - """ - Parses the hostnames element. - e.g. - - - """ - return [ - cls._parse_hostname(hname) for hname in hostnames_el.findall("hostname") - ] - - @classmethod - def _parse_hostname(cls, hostname_el: ET.Element) -> NmapHostname: - """ - Parses the hostname element. - e.g. - - :param hostname_el: XML tag from a nmap scan - """ - return NmapHostname.model_validate(dict(hostname_el.attrib)) - - @classmethod - def _parse_xml_ports( - cls, ports_elem: ET.Element - ) -> tuple[list[NmapPort], NmapPortStats]: - """ - Parses the list of scanned services from a targeted host. - """ - ports: list[NmapPort] = [] - stats = NmapPortStats() - - # handle extraports first - for e in ports_elem.findall("extraports"): - state = PortState(e.attrib["state"]) - count = int(e.attrib["count"]) - - key = state.value.replace("|", "_") - setattr(stats, key, getattr(stats, key) + count) - - for port_elem in ports_elem.findall("port"): - port = cls._parse_xml_port(port_elem) - ports.append(port) - key = port.state.value.replace("|", "_") - setattr(stats, key, getattr(stats, key) + 1) - return ports, stats - - @classmethod - def _parse_xml_service(cls, service_elem: ET.Element) -> NmapService: - svc = { - "name": service_elem.attrib.get("name"), - "product": service_elem.attrib.get("product"), - "version": service_elem.attrib.get("version"), - "extrainfo": service_elem.attrib.get("extrainfo"), - "method": service_elem.attrib.get("method"), - "conf": ( - int(service_elem.attrib["conf"]) - if "conf" in service_elem.attrib - else None - ), - "cpe": [e.text.strip() for e in service_elem.findall("cpe")], - } - - return NmapService.model_validate(svc) - - @classmethod - def _parse_xml_script(cls, script_elem: ET.Element) -> NmapScript: - output = script_elem.attrib.get("output") - if output: - output = output.strip() - script = { - "id": script_elem.attrib["id"], - "output": output, - } - - elements: dict[str, Any] = {} - - # handle value - for elem in script_elem.findall(".//elem"): - key = elem.attrib.get("key") - if key: - elements[key.strip()] = elem.text.strip() - - script["elements"] = elements - return NmapScript.model_validate(script) - - @classmethod - def _parse_xml_port(cls, port_elem: ET.Element) -> NmapPort: - """ - - - -
- Username and password - 2 -
- - - """ - state_elem = port_elem.find("state") - - port = { - "port": int(port_elem.attrib["portid"]), - "protocol": port_elem.attrib["protocol"], - "state": PortState(state_elem.attrib["state"]), - "reason": ( - PortStateReason(state_elem.attrib["reason"]) - if "reason" in state_elem.attrib - else None - ), - "reason_ttl": ( - int(state_elem.attrib["reason_ttl"]) - if "reason_ttl" in state_elem.attrib - else None - ), - } - - service_elem = port_elem.find("service") - if service_elem is not None: - port["service"] = cls._parse_xml_service(service_elem) - - port["scripts"] = [] - for script_elem in port_elem.findall("script"): - port["scripts"].append(cls._parse_xml_script(script_elem)) - - return NmapPort.model_validate(port) - - @classmethod - def _parse_trace(cls, host_elem: ET.Element) -> NmapTrace | None: - trace_elem = host_elem.find("trace") - if trace_elem is None: - return None - - port_attr = trace_elem.attrib.get("port") - proto_attr = trace_elem.attrib.get("proto") - - hops: list[NmapTraceHop] = [] - - for hop_elem in trace_elem.findall("hop"): - ttl = hop_elem.attrib.get("ttl") - if ttl is None: - continue # ttl is required by the DTD but guard anyway - - rtt = hop_elem.attrib.get("rtt") - ipaddr = hop_elem.attrib.get("ipaddr") - host = hop_elem.attrib.get("host") - - hops.append( - NmapTraceHop( - ttl=int(ttl), - ipaddr=ipaddr, - rtt_ms=float(rtt) if rtt is not None else None, - host=host, - ) - ) - - return NmapTrace( - port=int(port_attr) if port_attr is not None else None, - protocol=IPProtocol(proto_attr) if proto_attr is not None else None, - hops=hops, - ) - - -def parse_nmap_xml(raw) -> NmapResult: - return NmapXmlParser.parse_xml(raw) diff --git a/generalresearch/models/network/nmap/result.py b/generalresearch/models/network/nmap/result.py deleted file mode 100644 index 57c2e8b..0000000 --- a/generalresearch/models/network/nmap/result.py +++ /dev/null @@ -1,436 +0,0 @@ -from __future__ import annotations - -import json -from datetime import timedelta -from enum import StrEnum -from functools import cached_property -from typing import TYPE_CHECKING, Any, Literal - -from pydantic import BaseModel, Field, computed_field - -from generalresearch.models.network.definitions import IPProtocol - -if TYPE_CHECKING: - from generalresearch.models.custom_types import AwareDatetimeISO, IPvAnyAddressStr - - -class PortState(StrEnum): - OPEN = "open" - CLOSED = "closed" - FILTERED = "filtered" - UNFILTERED = "unfiltered" - OPEN_FILTERED = "open|filtered" - CLOSED_FILTERED = "closed|filtered" - # Added by me, does not get returned. Used for book-keeping - NOT_SCANNED = "not_scanned" - - -class PortStateReason(StrEnum): - SYN_ACK = "syn-ack" - RESET = "reset" - CONN_REFUSED = "conn-refused" - NO_RESPONSE = "no-response" - SYN = "syn" - FIN = "fin" - - ICMP_NET_UNREACH = "net-unreach" - ICMP_HOST_UNREACH = "host-unreach" - ICMP_PROTO_UNREACH = "proto-unreach" - ICMP_PORT_UNREACH = "port-unreach" - - ADMIN_PROHIBITED = "admin-prohibited" - HOST_PROHIBITED = "host-prohibited" - NET_PROHIBITED = "net-prohibited" - - ECHO_REPLY = "echo-reply" - TIME_EXCEEDED = "time-exceeded" - - -class NmapScanType(StrEnum): - SYN = "syn" - CONNECT = "connect" - ACK = "ack" - WINDOW = "window" - MAIMON = "maimon" - FIN = "fin" - NULL = "null" - XMAS = "xmas" - UDP = "udp" - SCTP_INIT = "sctpinit" - SCTP_COOKIE_ECHO = "sctpcookieecho" - - -class NmapHostState(StrEnum): - UP = "up" - DOWN = "down" - UNKNOWN = "unknown" - - -class NmapHostStatusReason(StrEnum): - USER_SET = "user-set" - SYN_ACK = "syn-ack" - RESET = "reset" - ECHO_REPLY = "echo-reply" - ARP_RESPONSE = "arp-response" - NO_RESPONSE = "no-response" - NET_UNREACH = "net-unreach" - HOST_UNREACH = "host-unreach" - PROTO_UNREACH = "proto-unreach" - PORT_UNREACH = "port-unreach" - ADMIN_PROHIBITED = "admin-prohibited" - LOCALHOST_RESPONSE = "localhost-response" - - -class NmapOSClass(BaseModel): - vendor: str = None - osfamily: str = None - osgen: str | None = None - accuracy: int = None - cpe: list[str] | None = None - - -class NmapOSMatch(BaseModel): - name: str - accuracy: int - classes: list[NmapOSClass] = Field(default_factory=list) - - @property - def best_class(self) -> NmapOSClass | None: - if not self.classes: - return None - return max(self.classes, key=lambda m: m.accuracy) - - -class NmapScript(BaseModel): - """ - - """ - - id: str - output: str | None = None - elements: dict[str, Any] = Field(default_factory=dict) - - -class NmapService(BaseModel): - # - name: str | None = None - product: str | None = None - version: str | None = None - extrainfo: str | None = None - method: str | None = None - conf: int | None = None - cpe: list[str] = Field(default_factory=list) - - def model_dump_postgres(self): - d = self.model_dump(mode="json") - d["service_name"] = self.name - return d - - -class NmapPort(BaseModel): - port: int = Field() - protocol: IPProtocol = Field() - # Closed ports will not have a NmapPort record - state: PortState = Field() - reason: PortStateReason | None = Field(default=None) - reason_ttl: int | None = Field(default=None) - - service: NmapService | None = None - scripts: list[NmapScript] = Field(default_factory=list) - - def model_dump_postgres(self, run_id: int): - # Writes for the network_portscanport table - d = {"port_scan_id": run_id} - data = self.model_dump( - mode="json", - include={ - "port", - "state", - "reason", - "reason_ttl", - }, - ) - d.update(data) - d["protocol"] = self.protocol.to_number() - if self.service: - d.update(self.service.model_dump_postgres()) - return d - - -class NmapHostScript(BaseModel): - id: str = Field() - output: str | None = Field(default=None) - - -class NmapTraceHop(BaseModel): - """ - One hop observed during Nmap's traceroute. - - Example XML: - - """ - - ttl: int = Field() - - ipaddr: str | None = Field( - default=None, - description="IP address of the responding router or host", - ) - - rtt_ms: float | None = Field( - default=None, - description="Round-trip time in milliseconds for the probe reaching this hop.", - ) - - host: str | None = Field( - default=None, - description="Reverse DNS hostname for the hop if Nmap resolved one.", - ) - - -class NmapTrace(BaseModel): - """ - Traceroute information collected by Nmap. - - Nmap performs a single traceroute per host using probes matching the scan - type (typically TCP) directed at a chosen destination port. - - Example XML: - - - ... - - """ - - port: int | None = Field( - default=None, - description="Destination port used for traceroute probes (may be absent depending on scan type).", - ) - protocol: IPProtocol | None = Field( - default=None, - description="Transport protocol used for the traceroute probes (tcp, udp, etc.).", - ) - - hops: list[NmapTraceHop] = Field( - default_factory=list, - description="Ordered list of hops observed during the traceroute.", - ) - - @property - def destination(self) -> NmapTraceHop | None: - return self.hops[-1] if self.hops else None - - -class NmapHostname(BaseModel): - # - name: str - type: Literal["PTR", "user"] | None = None - - -class NmapPortStats(BaseModel): - """ - This is counts across all protocols scanned (tcp/udp) - """ - - open: int = 0 - closed: int = 0 - filtered: int = 0 - unfiltered: int = 0 - open_filtered: int = 0 - closed_filtered: int = 0 - - -class NmapScanInfo(BaseModel): - """ - We could have multiple protocols in one run. - - - """ - - type: NmapScanType = Field() - protocol: IPProtocol = Field() - num_services: int = Field() - services: str = Field() - - @cached_property - 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() - for part in self.services.split(","): - if "-" in part: - start, end = part.split("-", 1) - ports.update(range(int(start), int(end) + 1)) - else: - ports.add(int(part)) - return ports - - -class NmapResult(BaseModel): - """ - A Nmap Run. Expects that we've only scanned ONE host. - """ - - command_line: str = Field() - started_at: AwareDatetimeISO = Field() - version: str = Field() - xmloutputversion: str = Field() - - scan_infos: list[NmapScanInfo] = Field(min_length=1) - - # comes from - finished_at: AwareDatetimeISO | None = Field(default=None) - exit_status: Literal["success", "error"] | None = Field(default=None) - - ##### - # Everything below here is from within the *single* host we've scanned - ##### - - # - host_state: NmapHostState = Field() - host_state_reason: NmapHostStatusReason = Field() - host_state_reason_ttl: int | None = None - - #
- target_ip: IPvAnyAddressStr = Field() - - hostnames: list[NmapHostname] = Field() - - ports: list[NmapPort] = [] - port_stats: NmapPortStats = Field() - - # - uptime_seconds: int | None = Field(default=None) - # - distance: int | None = Field(description="approx number of hops", default=None) - - # - tcp_sequence_index: int | None = None - tcp_sequence_difficulty: str | None = None - - # - ipid_sequence_class: str | None = None - - # - tcp_timestamp_class: str | None = None - - # - srtt_us: int | None = Field( - default=None, description="smoothed RTT estimate (microseconds µs)" - ) - rttvar_us: int | None = Field( - default=None, description="RTT variance (microseconds µs)" - ) - timeout_us: int | None = Field( - default=None, description="probe timeout (microseconds µs)" - ) - - os_matches: list[NmapOSMatch] | None = Field(default=None) - - host_scripts: list[NmapHostScript] = Field(default_factory=list) - - trace: NmapTrace | None = Field(default=None) - - raw_xml: str | None = None - - @computed_field - @property - def last_boot(self) -> AwareDatetimeISO | None: - if self.uptime_seconds: - return self.started_at - timedelta(seconds=self.uptime_seconds) - - @property - def scan_info_tcp(self): - return next( - filter(lambda x: x.protocol == IPProtocol.TCP, self.scan_infos), None - ) - - @property - def scan_info_udp(self): - return next( - filter(lambda x: x.protocol == IPProtocol.UDP, self.scan_infos), None - ) - - @property - def latency_ms(self) -> float | None: - return self.srtt_us / 1000 if self.srtt_us is not None else None - - @property - def best_os_match(self) -> NmapOSMatch | None: - if not self.os_matches: - return None - return max(self.os_matches, key=lambda m: m.accuracy) - - def filter_ports(self, protocol: IPProtocol, state: PortState) -> list[NmapPort]: - return [p for p in self.ports if p.protocol == protocol and p.state == state] - - @property - def tcp_open_ports(self) -> list[int]: - """ - Returns a list of open TCP port numbers. - """ - return [ - p.port - for p in self.filter_ports(protocol=IPProtocol.TCP, state=PortState.OPEN) - ] - - @property - def udp_open_ports(self) -> list[int]: - """ - Returns a list of open UDP port numbers. - """ - return [ - p.port - for p in self.filter_ports(protocol=IPProtocol.UDP, state=PortState.OPEN) - ] - - @cached_property - def _port_index(self) -> dict[tuple[IPProtocol, int], NmapPort]: - return {(p.protocol, p.port): p for p in self.ports} - - def get_port_state( - self, port: int, protocol: IPProtocol = IPProtocol.TCP - ) -> PortState: - # Explicit (only if scanned and not closed) - if (protocol, port) in self._port_index: - return self._port_index[(protocol, port)].state - - # Check if we even scanned it - scaninfo = next((s for s in self.scan_infos if s.protocol == protocol), None) - if scaninfo and port in scaninfo.port_set: - return PortState.CLOSED - - # We didn't scan it - return PortState.NOT_SCANNED - - def model_dump_postgres(self): - # Writes for the network_portscan table - d = {} - data = self.model_dump( - mode="json", - include={ - "started_at", - "host_state", - "host_state_reason", - "distance", - "uptime_seconds", - "raw_xml", - }, - ) - d.update(data) - d["ip"] = self.target_ip - d["xml_version"] = self.xmloutputversion - d["latency_ms"] = self.latency_ms - d["last_boot"] = self.last_boot - d["parsed"] = self.model_dump_json(indent=0) - d["open_tcp_ports"] = json.dumps(self.tcp_open_ports) - d["open_udp_ports"] = json.dumps(self.udp_open_ports) - return d diff --git a/generalresearch/models/network/rdns/__init__.py b/generalresearch/models/network/rdns/__init__.py deleted file mode 100644 index e69de29..0000000 diff --git a/generalresearch/models/network/rdns/command.py b/generalresearch/models/network/rdns/command.py deleted file mode 100644 index 2250449..0000000 --- a/generalresearch/models/network/rdns/command.py +++ /dev/null @@ -1,38 +0,0 @@ -import subprocess -from typing import TYPE_CHECKING - -from generalresearch.models.network.rdns.parser import parse_rdns_output - -if TYPE_CHECKING: - from generalresearch.models.network.rdns.result import RDNSResult - from generalresearch.models.network.tool_run_command import RDNSRunCommand - - -def run_rdns(config: RDNSRunCommand) -> RDNSResult: - cmd = config.to_command_str() - args = cmd.split(" ") - proc = subprocess.run( - args, - capture_output=True, - text=True, - check=False, - ) - raw = proc.stdout.strip() - return parse_rdns_output(ip=config.options.ip, raw=raw) - - -def build_rdns_command(ip: str) -> str: - # e.g. dig +noall +answer -x 1.2.3.4 - return f"dig +noall +answer -x {ip}" - - -def get_dig_version() -> str: - proc = subprocess.run( - ["dig", "-v"], - capture_output=True, - text=True, - check=False, - ) - # e.g. DiG 9.18.39-0ubuntu0.22.04.2-Ubuntu - ver_str = proc.stderr.strip() + proc.stdout.strip() - return ver_str.split("-", 1)[0].split(" ", 1)[1] diff --git a/generalresearch/models/network/rdns/execute.py b/generalresearch/models/network/rdns/execute.py deleted file mode 100644 index d6de84b..0000000 --- a/generalresearch/models/network/rdns/execute.py +++ /dev/null @@ -1,44 +0,0 @@ -from __future__ import annotations - -from datetime import UTC, datetime -from uuid import uuid4 - -from generalresearch.models.custom_types import UUIDStr -from generalresearch.models.network.rdns.command import ( - get_dig_version, - run_rdns, -) -from generalresearch.models.network.tool_run import ( - RDNSRun, - Status, - ToolClass, - ToolName, -) -from generalresearch.models.network.tool_run_command import ( - RDNSRunCommand, - RDNSRunCommandOptions, -) - - -def execute_rdns(ip: str, scan_group_id: UUIDStr | None = None): - 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=UTC) - - run = RDNSRun( - tool_name=ToolName.DIG, - tool_class=ToolClass.RDNS, - tool_version=tool_version, - status=Status.SUCCESS, - ip=ip, - started_at=started_at, - finished_at=finished_at, - raw_command=config.to_command_str(), - scan_group_id=scan_group_id or uuid4().hex, - config=config, - parsed=result, - ) - - return run diff --git a/generalresearch/models/network/rdns/parser.py b/generalresearch/models/network/rdns/parser.py deleted file mode 100644 index 31a5ed6..0000000 --- a/generalresearch/models/network/rdns/parser.py +++ /dev/null @@ -1,21 +0,0 @@ -import ipaddress -import re - -from generalresearch.models.custom_types import IPvAnyAddressStr -from generalresearch.models.network.rdns.result import RDNSResult - -PTR_RE = re.compile(r"\sPTR\s+([^\s]+)\.") - - -def parse_rdns_output(ip: IPvAnyAddressStr, raw: str) -> RDNSResult: - hostnames: list[str] = [] - - for line in raw.splitlines(): - m = PTR_RE.search(line) - if m: - hostnames.append(m.group(1)) - - return RDNSResult( - ip=ipaddress.ip_address(ip), - hostnames=hostnames, - ) diff --git a/generalresearch/models/network/rdns/result.py b/generalresearch/models/network/rdns/result.py deleted file mode 100644 index 46af643..0000000 --- a/generalresearch/models/network/rdns/result.py +++ /dev/null @@ -1,52 +0,0 @@ -from __future__ import annotations - -import json -from functools import cached_property - -import tldextract -from pydantic import BaseModel, Field, computed_field, model_validator - -from generalresearch.models.custom_types import IPvAnyAddressStr - - -class RDNSResult(BaseModel): - - ip: IPvAnyAddressStr = Field() - - hostnames: list[str] = Field(default_factory=list) - - @model_validator(mode="after") - def validate_hostname_prop(self): - assert len(self.hostnames) == self.hostname_count - if self.hostnames: - assert self.hostnames[0] == self.primary_hostname - assert self.primary_domain in self.primary_hostname - return self - - @computed_field(examples=["fixed-187-191-8-145.totalplay.net"]) - @cached_property - def primary_hostname(self) -> str | None: - if self.hostnames: - return self.hostnames[0] - - @computed_field(examples=[1]) - @cached_property - def hostname_count(self) -> int: - return len(self.hostnames) - - @computed_field(examples=["totalplay.net"]) - @cached_property - def primary_domain(self) -> str | None: - if self.primary_hostname: - return tldextract.extract( - self.primary_hostname - ).top_domain_under_public_suffix - - def model_dump_postgres(self): - # Writes for the network_rdnsresult table - d = self.model_dump( - mode="json", - include={"primary_hostname", "primary_domain", "hostname_count"}, - ) - d["hostnames"] = json.dumps(self.hostnames) - return d diff --git a/generalresearch/models/network/tool_run.py b/generalresearch/models/network/tool_run.py deleted file mode 100644 index 9088fe3..0000000 --- a/generalresearch/models/network/tool_run.py +++ /dev/null @@ -1,121 +0,0 @@ -from __future__ import annotations - -from enum import StrEnum -from typing import TYPE_CHECKING, Literal -from uuid import uuid4 - -from pydantic import BaseModel, Field, PositiveInt - -from generalresearch.models.custom_types import ( - AwareDatetimeISO, - IPvAnyAddressStr, - UUIDStr, -) - -if TYPE_CHECKING: - - from generalresearch.models.network.mtr.result import MTRResult - from generalresearch.models.network.nmap.result import NmapResult - from generalresearch.models.network.rdns.result import RDNSResult - from generalresearch.models.network.tool_run_command import ( - MTRRunCommand, - NmapRunCommand, - RDNSRunCommand, - ToolRunCommand, - ) - - -class ToolClass(StrEnum): - PORT_SCAN = "port_scan" - RDNS = "rdns" - PING = "ping" - TRACEROUTE = "traceroute" - - -class ToolName(StrEnum): - NMAP = "nmap" - RUSTMAP = "rustmap" - DIG = "dig" - PING = "ping" - TRACEROUTE = "traceroute" - MTR = "mtr" - - -class Status(StrEnum): - SUCCESS = "success" - FAILED = "failed" - TIMEOUT = "timeout" - ERROR = "error" - - -class ToolRun(BaseModel): - """ - A run of a networking tool against one host/ip. - """ - - id: PositiveInt | None = Field(default=None) - - ip: IPvAnyAddressStr = Field() - scan_group_id: UUIDStr = Field(default_factory=lambda: uuid4().hex) - tool_class: ToolClass = Field() - tool_name: ToolName = Field() - tool_version: str = Field() - - started_at: AwareDatetimeISO = Field() - finished_at: AwareDatetimeISO | None = Field(default=None) - status: Status | None = Field(default=None) - - raw_command: str = Field() - - config: ToolRunCommand = Field() - - def model_dump_postgres(self): - d = self.model_dump(mode="json", exclude={"config"}) - d["config"] = self.config.model_dump_json() - return d - - -class NmapRun(ToolRun): - tool_class: Literal[ToolClass.PORT_SCAN] = Field(default=ToolClass.PORT_SCAN) - tool_name: Literal[ToolName.NMAP] = Field(default=ToolName.NMAP) - config: NmapRunCommand = Field() - - parsed: NmapResult = Field() - - def model_dump_postgres(self): - d = super().model_dump_postgres() - d["run_id"] = self.id - d.update(self.parsed.model_dump_postgres()) - return d - - -class RDNSRun(ToolRun): - tool_class: Literal[ToolClass.RDNS] = Field(default=ToolClass.RDNS) - tool_name: Literal[ToolName.DIG] = Field(default=ToolName.DIG) - config: RDNSRunCommand = Field() - - parsed: RDNSResult = Field() - - def model_dump_postgres(self): - d = super().model_dump_postgres() - d["run_id"] = self.id - d.update(self.parsed.model_dump_postgres()) - return d - - -class MTRRun(ToolRun): - tool_class: Literal[ToolClass.TRACEROUTE] = Field(default=ToolClass.TRACEROUTE) - tool_name: Literal[ToolName.MTR] = Field(default=ToolName.MTR) - config: MTRRunCommand = Field() - - facility_id: int = Field(default=1) - source_ip: IPvAnyAddressStr = Field() - parsed: MTRResult = Field() - - def model_dump_postgres(self): - d = super().model_dump_postgres() - d["run_id"] = self.id - d["source_ip"] = self.source_ip - d["facility_id"] = self.facility_id - d.update(self.parsed.model_dump_postgres()) - return d diff --git a/generalresearch/models/network/tool_run_command.py b/generalresearch/models/network/tool_run_command.py deleted file mode 100644 index b07b811..0000000 --- a/generalresearch/models/network/tool_run_command.py +++ /dev/null @@ -1,68 +0,0 @@ -from __future__ import annotations - -from typing import TYPE_CHECKING, Literal - -from pydantic import BaseModel, Field - -from generalresearch.models.network.definitions import IPProtocol - -if TYPE_CHECKING: - from generalresearch.models.custom_types import IPvAnyAddressStr - - -class ToolRunCommand(BaseModel): - command: str = Field() - options: dict[str, str | int | None] = Field(default_factory=dict) - - -class NmapRunCommandOptions(BaseModel): - ip: IPvAnyAddressStr - top_ports: int | None = Field(default=1000) - ports: str | None = Field(default=None) - no_ping: bool = Field(default=True) - enable_advanced: bool = Field(default=True) - timing: int = Field(default=4) - - -class NmapRunCommand(ToolRunCommand): - command: Literal["nmap"] = Field(default="nmap") - options: NmapRunCommandOptions = Field() - - def to_command_str(self): - from generalresearch.models.network.nmap.command import build_nmap_command - - options = self.options - return build_nmap_command(**options.model_dump()) - - -class RDNSRunCommandOptions(BaseModel): - ip: IPvAnyAddressStr - - -class RDNSRunCommand(ToolRunCommand): - command: Literal["dig"] = Field(default="dig") - options: RDNSRunCommandOptions = Field() - - def to_command_str(self): - from generalresearch.models.network.rdns.command import build_rdns_command - - options = self.options - return build_rdns_command(**options.model_dump()) - - -class MTRRunCommandOptions(BaseModel): - ip: IPvAnyAddressStr = Field() - protocol: IPProtocol = Field(default=IPProtocol.ICMP) - port: int | None = Field(default=None) - report_cycles: int = Field(default=10) - - -class MTRRunCommand(ToolRunCommand): - command: Literal["mtr"] = Field(default="mtr") - options: MTRRunCommandOptions = Field() - - def to_command_str(self): - from generalresearch.models.network.mtr.command import build_mtr_command - - options = self.options - return build_mtr_command(**options.model_dump()) diff --git a/generalresearch/models/network/utils.py b/generalresearch/models/network/utils.py deleted file mode 100644 index fee9b80..0000000 --- a/generalresearch/models/network/utils.py +++ /dev/null @@ -1,5 +0,0 @@ -import requests - - -def get_source_ip(): - return requests.get("https://icanhazip.com?").text.strip() diff --git a/generalresearch/models/pollfish/question.py b/generalresearch/models/pollfish/question.py index f0c733c..aadf71e 100644 --- a/generalresearch/models/pollfish/question.py +++ b/generalresearch/models/pollfish/question.py @@ -4,17 +4,15 @@ from __future__ import annotations import json import logging from enum import StrEnum -from typing import TYPE_CHECKING, Any, Literal, Self +from typing import Any, Literal, Self from pydantic import BaseModel, Field, model_validator from generalresearch.models.definitions import Source from generalresearch.models.thl.profiling.marketplace import MarketplaceQuestion - -if TYPE_CHECKING: - from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestion, - ) +from generalresearch.models.thl.profiling.upk_question import ( + UpkQuestion, +) logging.basicConfig() logger = logging.getLogger() diff --git a/generalresearch/models/precision/question.py b/generalresearch/models/precision/question.py index 3e39124..97f6ca1 100644 --- a/generalresearch/models/precision/question.py +++ b/generalresearch/models/precision/question.py @@ -4,22 +4,20 @@ from __future__ import annotations import json import logging from enum import StrEnum -from typing import TYPE_CHECKING, Any, Literal +from typing import Any, Literal from pydantic import BaseModel, Field, ValidationError, field_validator, model_validator from generalresearch.models.definitions import Source +from generalresearch.models.precision import PrecisionQuestionID from generalresearch.models.string_utils import remove_nbsp from generalresearch.models.thl.profiling.marketplace import ( MarketplaceQuestion, MarketplaceUserQuestionAnswer, ) - -if TYPE_CHECKING: - from generalresearch.models.precision import PrecisionQuestionID - from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestion, - ) +from generalresearch.models.thl.profiling.upk_question import ( + UpkQuestion, +) logging.basicConfig() logger = logging.getLogger() diff --git a/generalresearch/models/precision/survey.py b/generalresearch/models/precision/survey.py index b77a365..cebe155 100644 --- a/generalresearch/models/precision/survey.py +++ b/generalresearch/models/precision/survey.py @@ -3,7 +3,7 @@ from __future__ import annotations import json from datetime import UTC from functools import cached_property -from typing import TYPE_CHECKING, Annotated, Any, Literal, Self +from typing import Annotated, Any, Literal, Self from more_itertools import flatten from pydantic import ( @@ -23,7 +23,7 @@ from generalresearch.models.custom_types import ( UUIDStrCoerce, ) from generalresearch.models.definitions import Source -from generalresearch.models.precision import PrecisionStatus +from generalresearch.models.precision import PrecisionQuestionID, PrecisionStatus from generalresearch.models.thl.demographics import Gender from generalresearch.models.thl.survey import MarketplaceTask from generalresearch.models.thl.survey.condition import ( @@ -31,9 +31,6 @@ from generalresearch.models.thl.survey.condition import ( MarketplaceCondition, ) -if TYPE_CHECKING: - from generalresearch.models.precision import PrecisionQuestionID - class PrecisionCondition(MarketplaceCondition): question_id: PrecisionQuestionID | None = Field() diff --git a/generalresearch/models/precision/task_collection.py b/generalresearch/models/precision/task_collection.py index 71241ad..c8db2af 100644 --- a/generalresearch/models/precision/task_collection.py +++ b/generalresearch/models/precision/task_collection.py @@ -1,18 +1,16 @@ -from typing import TYPE_CHECKING, Any +from typing import Any import pandas as pd from pandera.pandas import Check, Column, DataFrameSchema, Index from generalresearch.locales import Localelator from generalresearch.models.precision import PrecisionStatus +from generalresearch.models.precision.survey import PrecisionSurvey from generalresearch.models.thl.survey.task_collection import ( TaskCollection, create_empty_df_from_schema, ) -if TYPE_CHECKING: - from generalresearch.models.precision.survey import PrecisionSurvey - COUNTRY_ISOS = Localelator().get_all_countries() LANGUAGE_ISOS = Localelator().get_all_languages() diff --git a/generalresearch/models/prodege/question.py b/generalresearch/models/prodege/question.py index b963785..c0160bd 100644 --- a/generalresearch/models/prodege/question.py +++ b/generalresearch/models/prodege/question.py @@ -6,7 +6,7 @@ import logging from datetime import UTC, datetime from enum import StrEnum from functools import cached_property -from typing import TYPE_CHECKING, Any, Literal +from typing import Any, Literal from pydantic import ( BaseModel, @@ -18,15 +18,13 @@ from pydantic import ( ) from generalresearch.locales import Localelator +from generalresearch.models.custom_types import AwareDatetimeISO from generalresearch.models.definitions import MAX_INT32, Source +from generalresearch.models.prodege import ProdegeQuestionIdType from generalresearch.models.thl.profiling.marketplace import MarketplaceQuestion - -if TYPE_CHECKING: - from generalresearch.models.custom_types import AwareDatetimeISO - from generalresearch.models.prodege import ProdegeQuestionIdType - from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestion, - ) +from generalresearch.models.thl.profiling.upk_question import ( + UpkQuestion, +) logging.basicConfig() logger = logging.getLogger() diff --git a/generalresearch/models/prodege/survey.py b/generalresearch/models/prodege/survey.py index 26898d0..27034b0 100644 --- a/generalresearch/models/prodege/survey.py +++ b/generalresearch/models/prodege/survey.py @@ -7,7 +7,7 @@ from collections import defaultdict from datetime import UTC, datetime from decimal import Decimal from functools import cached_property -from typing import TYPE_CHECKING, Any, Literal +from typing import Any, Literal from pydantic import ( BaseModel, @@ -34,7 +34,9 @@ from generalresearch.models.definitions import ( ) from generalresearch.models.prodege import ( ProdegePastParticipationType, + ProdegeQuestionIdType, ProdegeStatus, + ProdgeRedirectStatus, ) from generalresearch.models.prodege.definitions import PG_COUNTRY_TO_ISO from generalresearch.models.thl.demographics import Gender @@ -44,13 +46,6 @@ from generalresearch.models.thl.survey.condition import ( MarketplaceCondition, ) -if TYPE_CHECKING: - - from generalresearch.models.prodege import ( - ProdegeQuestionIdType, - ProdgeRedirectStatus, - ) - logging.basicConfig() logger = logging.getLogger() logger.setLevel(logging.INFO) diff --git a/generalresearch/models/prodege/task_collection.py b/generalresearch/models/prodege/task_collection.py index 774fc7b..9f6a81b 100644 --- a/generalresearch/models/prodege/task_collection.py +++ b/generalresearch/models/prodege/task_collection.py @@ -1,18 +1,16 @@ -from typing import TYPE_CHECKING, Any +from typing import Any import pandas as pd from pandera.pandas import Check, Column, DataFrameSchema, Index from generalresearch.locales import Localelator from generalresearch.models.prodege import ProdegeStatus +from generalresearch.models.prodege.survey import ProdegeSurvey from generalresearch.models.thl.survey.task_collection import ( TaskCollection, create_empty_df_from_schema, ) -if TYPE_CHECKING: - from generalresearch.models.prodege.survey import ProdegeSurvey - COUNTRY_ISOS = Localelator().get_all_countries() LANGUAGE_ISOS = Localelator().get_all_languages() diff --git a/generalresearch/models/repdata/question.py b/generalresearch/models/repdata/question.py index 0ec102b..5115426 100644 --- a/generalresearch/models/repdata/question.py +++ b/generalresearch/models/repdata/question.py @@ -4,7 +4,7 @@ import json import logging from enum import StrEnum from functools import cached_property -from typing import TYPE_CHECKING, Any, Literal +from typing import Any, Literal from uuid import UUID from pydantic import ( @@ -20,11 +20,9 @@ from pydantic import ( from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.definitions import MAX_INT32, Source from generalresearch.models.thl.profiling.marketplace import MarketplaceQuestion - -if TYPE_CHECKING: - from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestion, - ) +from generalresearch.models.thl.profiling.upk_question import ( + UpkQuestion, +) logging.basicConfig() logger = logging.getLogger() diff --git a/generalresearch/models/repdata/task_collection.py b/generalresearch/models/repdata/task_collection.py index f2cb63b..aa591a9 100644 --- a/generalresearch/models/repdata/task_collection.py +++ b/generalresearch/models/repdata/task_collection.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Any +from typing import Any import pandas as pd from pandera.pandas import Check, Column, DataFrameSchema, Index @@ -8,14 +8,12 @@ from pandera.pandas import Check, Column, DataFrameSchema, Index from generalresearch.locales import Localelator from generalresearch.models.definitions import TaskCalculationType from generalresearch.models.repdata import RepDataStatus +from generalresearch.models.repdata.survey import RepDataSurveyHashed from generalresearch.models.thl.survey.task_collection import ( TaskCollection, create_empty_df_from_schema, ) -if TYPE_CHECKING: - from generalresearch.models.repdata.survey import RepDataSurveyHashed - COUNTRY_ISOS = Localelator().get_all_countries() LANGUAGE_ISOS = Localelator().get_all_languages() diff --git a/generalresearch/models/sago/question.py b/generalresearch/models/sago/question.py index 216b278..148b015 100644 --- a/generalresearch/models/sago/question.py +++ b/generalresearch/models/sago/question.py @@ -6,7 +6,7 @@ import json import logging from enum import StrEnum from functools import cached_property -from typing import TYPE_CHECKING, Any, Literal +from typing import Any, Literal from pydantic import ( BaseModel, @@ -18,15 +18,13 @@ from pydantic import ( model_validator, ) +from generalresearch.models.custom_types import AwareDatetimeISO from generalresearch.models.definitions import MAX_INT32, Source from generalresearch.models.string_utils import remove_nbsp from generalresearch.models.thl.profiling.marketplace import MarketplaceQuestion - -if TYPE_CHECKING: - from generalresearch.models.custom_types import AwareDatetimeISO - from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestion, - ) +from generalresearch.models.thl.profiling.upk_question import ( + UpkQuestion, +) logging.basicConfig() logger = logging.getLogger() diff --git a/generalresearch/models/sago/survey.py b/generalresearch/models/sago/survey.py index c9bf431..5330ace 100644 --- a/generalresearch/models/sago/survey.py +++ b/generalresearch/models/sago/survey.py @@ -5,7 +5,7 @@ import logging from datetime import UTC from decimal import Decimal from functools import cached_property -from typing import TYPE_CHECKING, Annotated, Any, Literal, Self +from typing import Annotated, Any, Literal, Self from more_itertools import flatten from pydantic import ( @@ -18,6 +18,14 @@ from pydantic import ( ) from generalresearch.locales import Localelator +from generalresearch.models.custom_types import ( + AlphaNumStr, + AlphaNumStrSet, + AwareDatetimeISO, + CoercedStr, + DeviceTypes, + IPLikeStrSet, +) from generalresearch.models.definitions import LogicalOperator, Source from generalresearch.models.sago import SagoStatus from generalresearch.models.thl.demographics import Gender @@ -27,16 +35,6 @@ from generalresearch.models.thl.survey.condition import ( MarketplaceCondition, ) -if TYPE_CHECKING: - from generalresearch.models.custom_types import ( - AlphaNumStr, - AlphaNumStrSet, - AwareDatetimeISO, - CoercedStr, - DeviceTypes, - IPLikeStrSet, - ) - logging.basicConfig() logger = logging.getLogger() logger.setLevel(logging.INFO) diff --git a/generalresearch/models/sago/task_collection.py b/generalresearch/models/sago/task_collection.py index 490f7a0..2879d9c 100644 --- a/generalresearch/models/sago/task_collection.py +++ b/generalresearch/models/sago/task_collection.py @@ -1,20 +1,18 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Any +from typing import Any import pandas as pd from pandera.pandas import Check, Column, DataFrameSchema, Index from generalresearch.locales import Localelator from generalresearch.models.sago import SagoStatus +from generalresearch.models.sago.survey import SagoSurvey from generalresearch.models.thl.survey.task_collection import ( TaskCollection, create_empty_df_from_schema, ) -if TYPE_CHECKING: - from generalresearch.models.sago.survey import SagoSurvey - COUNTRY_ISOS: set[str] = Localelator().get_all_countries() LANGUAGE_ISOS: set[str] = Localelator().get_all_languages() diff --git a/generalresearch/models/spectrum/question.py b/generalresearch/models/spectrum/question.py index 4b854fb..839f7a8 100644 --- a/generalresearch/models/spectrum/question.py +++ b/generalresearch/models/spectrum/question.py @@ -6,7 +6,7 @@ import logging from datetime import UTC, datetime from enum import IntEnum, StrEnum from functools import cached_property -from typing import TYPE_CHECKING, Any, Literal, Self +from typing import Any, Literal, Self from uuid import UUID from pydantic import ( @@ -18,18 +18,16 @@ from pydantic import ( model_validator, ) +from generalresearch.models.custom_types import AwareDatetimeISO from generalresearch.models.definitions import MAX_INT32, Source +from generalresearch.models.spectrum import SpectrumQuestionIdType from generalresearch.models.string_utils import remove_nbsp from generalresearch.models.thl.profiling.marketplace import ( MarketplaceQuestion, ) - -if TYPE_CHECKING: - from generalresearch.models.custom_types import AwareDatetimeISO - from generalresearch.models.spectrum import SpectrumQuestionIdType - from generalresearch.models.thl.profiling.upk_question import ( - UpkQuestion, - ) +from generalresearch.models.thl.profiling.upk_question import ( + UpkQuestion, +) logging.basicConfig() logger = logging.getLogger() diff --git a/generalresearch/models/spectrum/survey.py b/generalresearch/models/spectrum/survey.py index 6689d27..5b330a8 100644 --- a/generalresearch/models/spectrum/survey.py +++ b/generalresearch/models/spectrum/survey.py @@ -55,7 +55,7 @@ class SpectrumCondition(MarketplaceCondition): try: values = [tuple(map(int, v.split("-"))) for v in self.values] assert all(len(x) == 2 for x in values) - except (ValueError, AssertionError): + except ValueError, AssertionError: return self self.values = sorted( {str(val) for tupl in values for val in range(tupl[0], tupl[1] + 1)} @@ -75,7 +75,7 @@ class SpectrumCondition(MarketplaceCondition): rs["from"] = round(rs["from"] / 12) rs["to"] = round(rs["to"] / 12) d["values"] = [ - f"{rs["from"] or "inf"}-{rs["to"] or "inf"}" for rs in d["range_sets"] + f"{rs['from'] or 'inf'}-{rs['to'] or 'inf'}" for rs in d["range_sets"] ] d["value_type"] = ConditionValueType.RANGE return cls.model_validate(d) diff --git a/generalresearch/models/spectrum/task_collection.py b/generalresearch/models/spectrum/task_collection.py index 8e49434..609114e 100644 --- a/generalresearch/models/spectrum/task_collection.py +++ b/generalresearch/models/spectrum/task_collection.py @@ -1,21 +1,17 @@ from __future__ import annotations -from typing import TYPE_CHECKING - import pandas as pd from pandera.pandas import Check, Column, DataFrameSchema, Index from generalresearch.locales import Localelator from generalresearch.models.definitions import TaskCalculationType from generalresearch.models.spectrum import SpectrumStatus +from generalresearch.models.spectrum.survey import SpectrumSurvey from generalresearch.models.thl.survey.task_collection import ( TaskCollection, create_empty_df_from_schema, ) -if TYPE_CHECKING: - from generalresearch.models.spectrum.survey import SpectrumSurvey - COUNTRY_ISOS: set[str] = Localelator().get_all_countries() LANGUAGE_ISOS: set[str] = Localelator().get_all_languages() diff --git a/generalresearch/models/thl/contest/__init__.py b/generalresearch/models/thl/contest/__init__.py index c8342b3..0d7ace5 100644 --- a/generalresearch/models/thl/contest/__init__.py +++ b/generalresearch/models/thl/contest/__init__.py @@ -1,7 +1,7 @@ from __future__ import annotations from datetime import UTC, datetime -from typing import TYPE_CHECKING, Any, Self +from typing import Any, Self from uuid import uuid4 from pydantic import ( @@ -12,12 +12,10 @@ from pydantic import ( model_validator, ) +from generalresearch.currency import USDCent from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.thl.contest.definitions import ContestPrizeKind - -if TYPE_CHECKING: - from generalresearch.currency import USDCent - from generalresearch.models.thl.user import User +from generalresearch.models.thl.user import User class ContestEntryRule(BaseModel): @@ -88,9 +86,9 @@ class ContestPrize(BaseModel): @model_validator(mode="after") def validate_cash_value(self) -> Self: if self.kind == ContestPrizeKind.CASH: - assert ( - self.estimated_cash_value == self.cash_amount - ), "if kind is CASH, cash_amount must equal estimated_cash_value" + assert self.estimated_cash_value == self.cash_amount, ( + "if kind is CASH, cash_amount must equal estimated_cash_value" + ) return self diff --git a/generalresearch/models/thl/contest/contest.py b/generalresearch/models/thl/contest/contest.py index 5e30778..6fc60f6 100644 --- a/generalresearch/models/thl/contest/contest.py +++ b/generalresearch/models/thl/contest/contest.py @@ -3,7 +3,7 @@ from __future__ import annotations import json from abc import ABC, abstractmethod from datetime import UTC, datetime -from typing import TYPE_CHECKING, Any, Self +from typing import Any, Self from uuid import uuid4 from pydantic import ( @@ -19,18 +19,14 @@ from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.thl.contest import ( ContestEndCondition, ContestPrize, + ContestWinner, ) from generalresearch.models.thl.contest.definitions import ( ContestEndReason, ContestStatus, ContestType, ) - -if TYPE_CHECKING: - from generalresearch.models.thl.contest import ( - ContestWinner, - ) - from generalresearch.models.thl.locales import CountryISOs +from generalresearch.models.thl.locales import CountryISOs class ContestBase(BaseModel, ABC): diff --git a/generalresearch/models/thl/contest/contest_entry.py b/generalresearch/models/thl/contest/contest_entry.py index 261b3fc..4e90eb5 100644 --- a/generalresearch/models/thl/contest/contest_entry.py +++ b/generalresearch/models/thl/contest/contest_entry.py @@ -1,7 +1,7 @@ from __future__ import annotations from datetime import UTC, datetime -from typing import TYPE_CHECKING, Any +from typing import Any from uuid import uuid4 from pydantic import ( @@ -16,9 +16,7 @@ from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.thl.contest.definitions import ( ContestEntryType, ) - -if TYPE_CHECKING: - from generalresearch.models.thl.user import User +from generalresearch.models.thl.user import User class ContestEntryCreate(BaseModel): diff --git a/generalresearch/models/thl/contest/leaderboard.py b/generalresearch/models/thl/contest/leaderboard.py index efbbd0a..064c0d1 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 UTC, datetime, timedelta -from typing import TYPE_CHECKING, Any, Literal, Self +from typing import Any, Literal, Self from pydantic import ( ConfigDict, @@ -16,6 +16,9 @@ from generalresearch.currency import USDCent from generalresearch.decorators import LOG from generalresearch.managers.leaderboard import country_timezone from generalresearch.managers.leaderboard.manager import LeaderboardManager +from generalresearch.managers.thl.user_manager.user_manager import ( + UserManager, +) from generalresearch.models.thl.contest import ( ContestEndCondition, ContestPrize, @@ -39,11 +42,6 @@ from generalresearch.models.thl.leaderboard import ( LeaderboardFrequency, ) -if TYPE_CHECKING: - from generalresearch.managers.thl.user_manager.user_manager import ( - UserManager, - ) - class LeaderboardContestCreate(ContestBase): model_config = ConfigDict( @@ -75,9 +73,9 @@ class LeaderboardContestCreate(ContestBase): ranks = {x.leaderboard_rank for x in self.prizes} assert None not in ranks, "Must have leaderboard_rank defined" assert min(ranks) == 1, "Must start with rank 1" - assert ranks == set( - range(min(ranks), max(ranks) + 1) - ), "cannot skip prize leaderboard_ranks" + assert ranks == set(range(min(ranks), max(ranks) + 1)), ( + "cannot skip prize leaderboard_ranks" + ) return self @model_validator(mode="after") @@ -88,9 +86,9 @@ class LeaderboardContestCreate(ContestBase): @model_validator(mode="after") def check_end_condition(self) -> Self: - assert ( - not self.end_condition.target_entry_amount - ), "target_entry_amount not valid in leaderboard contest" + assert not self.end_condition.target_entry_amount, ( + "target_entry_amount not valid in leaderboard contest" + ) # the ends_at will get set automatically from the leaderboard_key return self @@ -174,13 +172,13 @@ class LeaderboardContest(LeaderboardContestCreate, Contest): @model_validator(mode="after") def validate_product_lb_key(self) -> Self: - assert ( - self.product_id == self.leaderboard_key_parts["product_id"] - ), "leaderboard_key product_id is invalid" + assert self.product_id == self.leaderboard_key_parts["product_id"], ( + "leaderboard_key product_id is invalid" + ) if self.country_isos: - assert ( - len(self.country_isos) == 1 - ), "Can only set 1 country_iso in a leaderboard contest" + assert len(self.country_isos) == 1, ( + "Can only set 1 country_iso in a leaderboard contest" + ) assert ( next(iter(self.country_isos)) == self.leaderboard_key_parts["country_iso"] @@ -192,9 +190,9 @@ class LeaderboardContest(LeaderboardContestCreate, Contest): @model_validator(mode="after") def validate_tie_break(self) -> Self: if self.tie_break_strategy == LeaderboardTieBreakStrategy.SPLIT_PRIZE_POOL: - assert all( - p.kind == ContestPrizeKind.CASH for p in self.prizes - ), "All prizes must be cash due to the tie-break strategy" + assert all(p.kind == ContestPrizeKind.CASH for p in self.prizes), ( + "All prizes must be cash due to the tie-break strategy" + ) return self @model_validator(mode="after") diff --git a/generalresearch/models/thl/contest/raffle.py b/generalresearch/models/thl/contest/raffle.py index 21bc481..7e84e89 100644 --- a/generalresearch/models/thl/contest/raffle.py +++ b/generalresearch/models/thl/contest/raffle.py @@ -4,7 +4,7 @@ import logging import random from collections import defaultdict from datetime import UTC, datetime -from typing import TYPE_CHECKING, Any, Literal, Self +from typing import Any, Literal, Self from pydantic import ( ConfigDict, @@ -20,12 +20,16 @@ from generalresearch.models.thl.contest import ( ContestEndCondition, ContestEntryRule, ContestPrize, + ContestWinner, ) from generalresearch.models.thl.contest.contest import ( Contest, ContestBase, ContestUserView, ) +from generalresearch.models.thl.contest.contest_entry import ( + ContestEntry, +) from generalresearch.models.thl.contest.definitions import ( ContestEndReason, ContestEntryType, @@ -34,14 +38,6 @@ from generalresearch.models.thl.contest.definitions import ( ContestType, ) -if TYPE_CHECKING: - from generalresearch.models.thl.contest import ( - ContestWinner, - ) - from generalresearch.models.thl.contest.contest_entry import ( - ContestEntry, - ) - logging.basicConfig() LOG = logging.getLogger() LOG.setLevel(logging.INFO) @@ -114,9 +110,9 @@ class RaffleContest(RaffleContestCreate, Contest): @model_validator(mode="after") def validate_entry_type(self): - assert all( - entry.entry_type == self.entry_type for entry in self.entries - ), f"all entries must be of type {self.entry_type}" + assert all(entry.entry_type == self.entry_type for entry in self.entries), ( + f"all entries must be of type {self.entry_type}" + ) return self @field_validator("current_amount", mode="before") diff --git a/generalresearch/models/thl/profiling/upk_property.py b/generalresearch/models/thl/profiling/upk_property.py index 96f1b4c..e46e00a 100644 --- a/generalresearch/models/thl/profiling/upk_property.py +++ b/generalresearch/models/thl/profiling/upk_property.py @@ -2,17 +2,14 @@ from __future__ import annotations from enum import StrEnum from functools import cached_property -from typing import TYPE_CHECKING from uuid import uuid4 from pydantic import BaseModel, ConfigDict, Field, TypeAdapter from generalresearch.models.custom_types import CountryISOLike, UUIDStr +from generalresearch.models.thl.category import Category from generalresearch.utils.enum import ReprEnumMeta -if TYPE_CHECKING: - from generalresearch.models.thl.category import Category - class PropertyType(StrEnum, metaclass=ReprEnumMeta): # UserProfileKnowledge Item diff --git a/generalresearch/models/thl/profiling/user_info.py b/generalresearch/models/thl/profiling/user_info.py index 5124d17..46733d4 100644 --- a/generalresearch/models/thl/profiling/user_info.py +++ b/generalresearch/models/thl/profiling/user_info.py @@ -1,18 +1,14 @@ from __future__ import annotations -from typing import TYPE_CHECKING - from pydantic import BaseModel, ConfigDict, Field from pydantic.json_schema import SkipJsonSchema from generalresearch.models.custom_types import AwareDatetimeISO - -if TYPE_CHECKING: - from generalresearch.models.definitions import Source - from generalresearch.models.thl.profiling.user_question_answer import ( - MarketplaceResearchProfileQuestion, - ) - from generalresearch.models.thl.user import User +from generalresearch.models.definitions import Source +from generalresearch.models.thl.profiling.user_question_answer import ( + MarketplaceResearchProfileQuestion, +) +from generalresearch.models.thl.user import User class UserProfileKnowledgeAnswer(BaseModel): diff --git a/generalresearch/models/thl/survey/__init__.py b/generalresearch/models/thl/survey/__init__.py index 76f819e..7749bdb 100644 --- a/generalresearch/models/thl/survey/__init__.py +++ b/generalresearch/models/thl/survey/__init__.py @@ -3,32 +3,27 @@ from __future__ import annotations from abc import ABC, abstractmethod from decimal import Decimal from itertools import product -from typing import TYPE_CHECKING from more_itertools import flatten from pydantic import BaseModel, Field +from generalresearch.models.definitions import Source from generalresearch.models.thl.demographics import ( AgeGroup, DemographicTarget, Gender, ) +from generalresearch.models.thl.locales import ( + CountryISO, + CountryISOs, + LanguageISO, + LanguageISOs, +) from generalresearch.models.thl.survey.condition import ( ConditionValueType, + MarketplaceCondition, ) -if TYPE_CHECKING: - from generalresearch.models.definitions import Source - from generalresearch.models.thl.locales import ( - CountryISO, - CountryISOs, - LanguageISO, - LanguageISOs, - ) - from generalresearch.models.thl.survey.condition import ( - MarketplaceCondition, - ) - class MarketplaceTask(BaseModel, ABC): """This is called a "Task" even though generally it represents a survey diff --git a/test_utils/managers/network/__init__.py b/test_utils/managers/network/__init__.py deleted file mode 100644 index e69de29..0000000 diff --git a/test_utils/managers/network/conftest.py b/test_utils/managers/network/conftest.py deleted file mode 100644 index e69de29..0000000 diff --git a/test_utils/models/network/__init__.py b/test_utils/models/network/__init__.py deleted file mode 100644 index e69de29..0000000 diff --git a/test_utils/models/network/conftest.py b/test_utils/models/network/conftest.py deleted file mode 100644 index c7bcc7e..0000000 --- a/test_utils/models/network/conftest.py +++ /dev/null @@ -1,145 +0,0 @@ -import os -from datetime import UTC, datetime, timedelta -from typing import TYPE_CHECKING -from uuid import uuid4 - -import pytest -from pytest import FixtureRequest as Request - -from generalresearch.managers.network.label import IPLabelManager -from generalresearch.managers.network.tool_run import ToolRunManager -from generalresearch.models.network.definitions import IPProtocol -from generalresearch.models.network.mtr.parser import parse_mtr_output -from generalresearch.models.network.mtr.result import MTRResult -from generalresearch.models.network.nmap.parser import parse_nmap_xml -from generalresearch.models.network.nmap.result import NmapResult -from generalresearch.models.network.rdns.parser import parse_rdns_output -from generalresearch.models.network.rdns.result import RDNSResult -from generalresearch.models.network.tool_run import MTRRun, NmapRun, RDNSRun, Status -from generalresearch.models.network.tool_run_command import ( - MTRRunCommand, - MTRRunCommandOptions, - NmapRunCommand, - NmapRunCommandOptions, - RDNSRunCommand, - RDNSRunCommandOptions, -) -from generalresearch.pg_helper import PostgresConfig - - -@pytest.fixture(scope="session") -def scan_group_id() -> str: - return uuid4().hex - - -@pytest.fixture(scope="session") -def iplabel_manager(thl_web_rw: PostgresConfig) -> IPLabelManager: - return IPLabelManager(pg_config=thl_web_rw) - - -@pytest.fixture(scope="session") -def toolrun_manager(thl_web_rw: PostgresConfig) -> ToolRunManager: - return ToolRunManager(pg_config=thl_web_rw) - - -@pytest.fixture(scope="session") -def nmap_raw_output(request: Request) -> str: - fp = os.path.join(request.config.rootpath, "data/nmaprun1.xml") - with open(fp) as f: - data = f.read() - return data - - -@pytest.fixture(scope="session") -def nmap_result(nmap_raw_output: str) -> NmapResult: - return parse_nmap_xml(nmap_raw_output) - - -@pytest.fixture(scope="session") -def nmap_run(nmap_result: NmapResult, scan_group_id: str): - r = nmap_result - config = NmapRunCommand( - command="nmap", - options=NmapRunCommandOptions( - ip=r.target_ip, ports="22-1000,11000,1100,3389,61232", top_ports=None - ), - ) - return NmapRun( - tool_version=r.version, - status=Status.SUCCESS, - ip=r.target_ip, - started_at=r.started_at, - finished_at=r.finished_at, - raw_command=config.to_command_str(), - scan_group_id=scan_group_id, - config=config, - parsed=r, - ) - - -@pytest.fixture(scope="session") -def dig_raw_output() -> str: - return "156.32.33.45.in-addr.arpa. 300 IN PTR scanme.nmap.org." - - -@pytest.fixture(scope="session") -def rdns_result(dig_raw_output: str) -> RDNSResult: - return parse_rdns_output(ip="45.33.32.156", raw=dig_raw_output) - - -@pytest.fixture(scope="session") -def rdns_run(rdns_result: RDNSResult, scan_group_id: str): - r = rdns_result - ip = "45.33.32.156" - utc_now = datetime.now(tz=UTC) - config = RDNSRunCommand(command="dig", options=RDNSRunCommandOptions(ip=ip)) - return RDNSRun( - tool_version="1.2.3", - status=Status.SUCCESS, - ip=ip, - started_at=utc_now, - finished_at=utc_now + timedelta(seconds=1), - raw_command=config.to_command_str(), - scan_group_id=scan_group_id, - config=config, - parsed=r, - ) - - -@pytest.fixture(scope="session") -def mtr_raw_output(request: Request) -> str: - fp = os.path.join(request.config.rootpath, "data/mtr_fatbeam.json") - with open(fp) as f: - data = f.read() - return data - - -@pytest.fixture(scope="session") -def mtr_result(mtr_raw_output: str) -> MTRResult: - return parse_mtr_output(mtr_raw_output, port=443, protocol=IPProtocol.TCP) - - -@pytest.fixture(scope="session") -def mtr_run(mtr_result: MTRResult, scan_group_id: str): - r = mtr_result - utc_now = datetime.now(tz=UTC) - config = MTRRunCommand( - command="mtr", - options=MTRRunCommandOptions( - ip=r.destination, protocol=IPProtocol.TCP, port=443 - ), - ) - - return MTRRun( - tool_version="1.2.3", - status=Status.SUCCESS, - ip=r.destination, - started_at=utc_now, - finished_at=utc_now + timedelta(seconds=1), - raw_command=config.to_command_str(), - scan_group_id=scan_group_id, - config=config, - parsed=r, - facility_id=1, - source_ip="1.2.3.4", - ) diff --git a/tests/managers/network/__init__.py b/tests/managers/network/__init__.py deleted file mode 100644 index e69de29..0000000 diff --git a/tests/managers/network/test_label.py b/tests/managers/network/test_label.py deleted file mode 100644 index abdd28f..0000000 --- a/tests/managers/network/test_label.py +++ /dev/null @@ -1,209 +0,0 @@ -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.models.network.label import ( - IPLabel, - IPLabelKind, - IPLabelMetadata, - IPLabelSource, -) -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: datetime) -> IPLabel: - ip = ipaddress.IPv6Network((fake.ipv6(), 64), strict=False) - return IPLabel( - label_kind=IPLabelKind.VPN, - labeled_at=utc_now, - source=IPLabelSource.INTERNAL_USE, - provider="GeoNodE", - created_at=utc_now, - ip=ip, - metadata=IPLabelMetadata(services=["RDP"]), - ) - - -def test_model(utc_now: datetime): - ip = fake.ipv4_public() - lbl = IPLabel( - label_kind=IPLabelKind.VPN, - labeled_at=utc_now, - source=IPLabelSource.INTERNAL_USE, - provider="GeoNodE", - created_at=utc_now, - ip=ip, - ) - assert lbl.ip.prefixlen == 32 - print(f"{lbl.ip=}") - - ip = ipaddress.IPv4Network((ip, 24), strict=False) - lbl = IPLabel( - label_kind=IPLabelKind.VPN, - labeled_at=utc_now, - source=IPLabelSource.INTERNAL_USE, - provider="GeoNodE", - created_at=utc_now, - ip=ip, - ) - print(f"{lbl.ip=}") - - with pytest.raises(ValidationError, match="IPv6 network must be /64 or larger"): - IPLabel( - label_kind=IPLabelKind.VPN, - labeled_at=utc_now, - source=IPLabelSource.INTERNAL_USE, - provider="GeoNodE", - created_at=utc_now, - ip=fake.ipv6(), - ) - - ip = ipaddress.IPv6Network((fake.ipv6(), 64), strict=False) - lbl = IPLabel( - label_kind=IPLabelKind.VPN, - labeled_at=utc_now, - source=IPLabelSource.INTERNAL_USE, - provider="GeoNodE", - created_at=utc_now, - ip=ip, - ) - print(f"{lbl.ip=}") - - ip = ipaddress.IPv6Network((ip.network_address, 48), strict=False) - lbl = IPLabel( - label_kind=IPLabelKind.VPN, - labeled_at=utc_now, - source=IPLabelSource.INTERNAL_USE, - provider="GeoNodE", - created_at=utc_now, - ip=ip, - ) - print(f"{lbl.ip=}") - - -def test_create(iplabel_manager: IPLabelManager, ip_label: IPLabel): - iplabel_manager.create(ip_label) - - with pytest.raises( - UniqueViolation, match="duplicate key value violates unique constraint" - ): - iplabel_manager.create(ip_label) - - -def test_filter(iplabel_manager: IPLabelManager, ip_label: IPLabel, utc_hour_ago): - res = iplabel_manager.filter(ips=[ip_label.ip]) - assert len(res) == 0 - - iplabel_manager.create(ip_label) - res = iplabel_manager.filter(ips=[ip_label.ip]) - assert len(res) == 1 - - out = res[0] - assert out == ip_label - - res = iplabel_manager.filter(ips=[ip_label.ip], labeled_after=utc_hour_ago) - assert len(res) == 1 - - ip_label2 = ip_label.model_copy() - ip_label2.ip = fake.ipv4_public() - iplabel_manager.create(ip_label2) - res = iplabel_manager.filter(ips=[ip_label.ip, ip_label2.ip]) - assert len(res) == 2 - - -def test_filter_network( - iplabel_manager: IPLabelManager, ip_label: IPLabel, utc_hour_ago -): - print(ip_label) - ip_label = ip_label.model_copy() - ip_label.ip = ipaddress.IPv6Network((fake.ipv6(), 64), strict=False) - - iplabel_manager.create(ip_label) - res = iplabel_manager.filter(ips=[ip_label.ip]) - assert len(res) == 1 - - out = res[0] - assert out == ip_label - - res = iplabel_manager.filter(ips=[ip_label.ip], labeled_after=utc_hour_ago) - assert len(res) == 1 - - ip_label2 = ip_label.model_copy() - ip_label2.ip = fake.ipv4_public() - iplabel_manager.create(ip_label2) - res = iplabel_manager.filter(ips=[ip_label.ip, ip_label2.ip]) - assert len(res) == 2 - - -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() - # Generally, we'd want to annotate the /64 network - # e.g. '51b7:b38d:8717:6c5b::/64' - ip_64 = ipaddress.IPv6Network((ip, 64), strict=False) - - label = IPLabel( - label_kind=IPLabelKind.VPN, - labeled_at=utc_now, - source=IPLabelSource.INTERNAL_USE, - provider="GeoNodE", - created_at=utc_now, - ip=ip_64, - ) - iplabel_manager.create(label) - - # If I query for the /128 directly, I won't find it - res = iplabel_manager.filter(ips=[ip]) - assert len(res) == 0 - - # If I query for the /64 network I will - res = iplabel_manager.filter(ips=[ip_64]) - assert len(res) == 1 - - # Or, I can query for the /128 ip IN a network - res = iplabel_manager.filter(ip_in_network=ip) - assert len(res) == 1 - - -def test_label_cidr_and_ipinfo( - 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 - ip = fake.ipv6() - ip_information_factory(ip=ip, geoname=ip_geoname) - # We normalize for storage into ipinfo table - ip_norm, _ = normalize_ip(ip) - - # Test with a larger network - ip_48 = ipaddress.IPv6Network((ip, 48), strict=False) - print(f"{ip=}") - print(f"{ip_norm=}") - print(f"{ip_48=}") - label = IPLabel( - label_kind=IPLabelKind.VPN, - labeled_at=utc_now, - source=IPLabelSource.INTERNAL_USE, - provider="GeoNodE", - created_at=utc_now, - ip=ip_48, - ) - iplabel_manager.create(label) - - res = iplabel_manager.test_join(ip_norm) - print(res) diff --git a/tests/managers/network/test_tool_run.py b/tests/managers/network/test_tool_run.py deleted file mode 100644 index a815809..0000000 --- a/tests/managers/network/test_tool_run.py +++ /dev/null @@ -1,25 +0,0 @@ -def test_create_tool_run_from_nmap_run(nmap_run, toolrun_manager): - - toolrun_manager.create_nmap_run(nmap_run) - - run_out = toolrun_manager.get_nmap_run(nmap_run.id) - - assert nmap_run == run_out - - -def test_create_tool_run_from_rdns_run(rdns_run, toolrun_manager): - - toolrun_manager.create_rdns_run(rdns_run) - - run_out = toolrun_manager.get_rdns_run(rdns_run.id) - - assert rdns_run == run_out - - -def test_create_tool_run_from_mtr_run(mtr_run, toolrun_manager): - - toolrun_manager.create_mtr_run(mtr_run) - - run_out = toolrun_manager.get_mtr_run(mtr_run.id) - - assert mtr_run == run_out diff --git a/tests/models/network/__init__.py b/tests/models/network/__init__.py deleted file mode 100644 index e69de29..0000000 diff --git a/tests/models/network/test_mtr.py b/tests/models/network/test_mtr.py deleted file mode 100644 index 5d136c4..0000000 --- a/tests/models/network/test_mtr.py +++ /dev/null @@ -1,33 +0,0 @@ -from __future__ import annotations - -from typing import TYPE_CHECKING - -import faker - -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() - - -def test_execute_mtr(toolrun_manager: ToolRunManager): - ip = "65.19.129.53" - - run = execute_mtr(ip=ip, report_cycles=3) - assert run.tool_name == ToolName.MTR - assert run.tool_class == ToolClass.TRACEROUTE - assert run.ip == ip - result = run.parsed - - last_hop = result.hops[-1] - assert last_hop.asn == 6939 - assert last_hop.domain == "grlengine.com" - - last_hop_1 = result.hops[-2] - assert last_hop_1.asn == 6939 - assert last_hop_1.domain == "he.net" - - toolrun_manager.create_mtr_run(run) diff --git a/tests/models/network/test_nmap.py b/tests/models/network/test_nmap.py deleted file mode 100644 index 6adc9e4..0000000 --- a/tests/models/network/test_nmap.py +++ /dev/null @@ -1,39 +0,0 @@ -from __future__ import annotations - -import subprocess -from typing import TYPE_CHECKING - -import faker - -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 ToolClass, ToolName - -if TYPE_CHECKING: - from generalresearch.managers.network.tool_run import ToolRunManager - from generalresearch.models.network.tool_run import NmapRun - -fake = faker.Faker() - - -def resolve(host: str): - return subprocess.check_output(["dig", host, "+short"]).decode().strip() - - -def test_execute_nmap_scanme(toolrun_manager: ToolRunManager): - ip = resolve("scanme.nmap.org") - - run: NmapRun = execute_nmap( - ip=ip, top_ports=None, ports="20-30", enable_advanced=False - ) - assert run.tool_name == ToolName.NMAP - assert run.tool_class == ToolClass.PORT_SCAN - assert run.ip == ip - assert isinstance(run.parsed, NmapResult) - result = run.parsed - - port22 = result._port_index[(IPProtocol.TCP, 22)] - assert port22.state == PortState.OPEN - - toolrun_manager.create_nmap_run(run) diff --git a/tests/models/network/test_nmap_parser.py b/tests/models/network/test_nmap_parser.py deleted file mode 100644 index fc9884b..0000000 --- a/tests/models/network/test_nmap_parser.py +++ /dev/null @@ -1,32 +0,0 @@ -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 NmapTrace - -if TYPE_CHECKING: - from generalresearch.models.network.nmap.result import NmapResult - - -@pytest.fixture -def nmap_raw_output_2(request) -> str: - fp = os.path.join(request.config.rootpath, "data/nmaprun2.xml") - with open(fp) as f: - data = f.read() - return data - - -def test_nmap_xml_parser(nmap_raw_output: str, nmap_raw_output_2: str): - n: NmapResult = parse_nmap_xml(nmap_raw_output) - assert n.tcp_open_ports == [61232] - - assert isinstance(n.trace, NmapTrace) - assert len(n.trace.hops) == 18 - - n = parse_nmap_xml(nmap_raw_output_2) - assert n.tcp_open_ports == [22, 80, 9929, 31337] - assert n.trace is None diff --git a/tests/models/network/test_rdns.py b/tests/models/network/test_rdns.py deleted file mode 100644 index 82126dd..0000000 --- a/tests/models/network/test_rdns.py +++ /dev/null @@ -1,40 +0,0 @@ -from __future__ import annotations - -from typing import TYPE_CHECKING - -import faker - -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() - - -def test_execute_rdns_grl(toolrun_manager: ToolRunManager): - ip = "65.19.129.53" - run = execute_rdns(ip=ip) - assert run.tool_name == ToolName.DIG - assert run.tool_class == ToolClass.RDNS - assert run.ip == ip - result = run.parsed - assert result.primary_hostname == "in1-smtp.grlengine.com" - assert result.primary_domain == "grlengine.com" - assert result.hostname_count == 1 - - toolrun_manager.create_rdns_run(run) - - -def test_execute_rdns_none(toolrun_manager: ToolRunManager): - ip = fake.ipv6() - run = execute_rdns(ip) - result = run.parsed - - assert result.primary_hostname is None - assert result.primary_domain is None - assert result.hostname_count == 0 - assert result.hostnames == [] - - toolrun_manager.create_rdns_run(run) -- cgit v1.2.3 From f8f1f07b193845d92c7f6ef8ae95b9696db6330f Mon Sep 17 00:00:00 2001 From: stuppie Date: Fri, 4 Sep 2026 12:54:13 -0600 Subject: fix a lot of tests --- generalresearch/currency.py | 12 ++--- generalresearch/grliq/models/forensic_data.py | 14 ++--- .../incite/mergers/foundations/enriched_wall.py | 3 -- pyproject.toml | 2 +- test_utils/managers/gr/conftest.py | 2 + test_utils/managers/thl/conftest.py | 5 ++ test_utils/models/gr/conftest.py | 14 ++--- tests/conftest.py | 2 - tests/models/custom_types/test_dsn.py | 5 +- tests/models/gr/test_authentication.py | 63 ++++++++++------------ tests/models/gr/test_business.py | 25 ++++----- tests/models/gr/test_team.py | 8 ++- tests/models/test_finance.py | 14 ++--- .../thl/test_contest/test_leaderboard_contest.py | 8 +-- tests/models/thl/test_payout_format.py | 8 +-- tests/models/thl/test_product.py | 21 ++------ 16 files changed, 84 insertions(+), 122 deletions(-) (limited to 'test_utils/models') diff --git a/generalresearch/currency.py b/generalresearch/currency.py index 716cb0f..7a9d037 100644 --- a/generalresearch/currency.py +++ b/generalresearch/currency.py @@ -29,12 +29,12 @@ class USDCent(int): if isinstance(value, float): warnings.warn( - "USDCent init with a float. Rounding behavior may " "be unexpected" + "USDCent init with a float. Rounding behavior may be unexpected" ) if isinstance(value, Decimal): warnings.warn( - "USDCent init with a Decimal. Rounding behavior may " "be unexpected" + "USDCent init with a Decimal. Rounding behavior may be unexpected" ) if value < 0: @@ -61,7 +61,7 @@ class USDCent(int): res = super().__abs__() return self.__class__(res) - def __truediv__(self): + def __truediv__(self, value): raise ValueError("Division not allowed for USDCent") def __str__(self): @@ -97,12 +97,12 @@ class USDMill(int): if isinstance(value, float): warnings.warn( - "USDMill init with a float. Rounding behavior " "may be unexpected" + "USDMill init with a float. Rounding behavior may be unexpected" ) if isinstance(value, Decimal): warnings.warn( - "USDMill init with a Decimal. Rounding behavior " "may be unexpected" + "USDMill init with a Decimal. Rounding behavior may be unexpected" ) if value < 0: @@ -129,7 +129,7 @@ class USDMill(int): res = super().__abs__() return self.__class__(res) - def __truediv__(self): + def __truediv__(self, value): raise ValueError("Division not allowed for USDMill") def __str__(self): diff --git a/generalresearch/grliq/models/forensic_data.py b/generalresearch/grliq/models/forensic_data.py index 6a07774..9d69e41 100644 --- a/generalresearch/grliq/models/forensic_data.py +++ b/generalresearch/grliq/models/forensic_data.py @@ -53,9 +53,9 @@ from generalresearch.models.custom_types import ( IPvAnyAddressStr, UUIDStr, ) +from generalresearch.models.thl.ipinfo import GeoIPInformation if TYPE_CHECKING: - from generalresearch.models.thl.ipinfo import GeoIPInformation from generalresearch.models.thl.session import Session fake = Faker() @@ -776,14 +776,14 @@ class GrlIqData(BaseModel): # product_id and product_user_id are parsed from the post body. make sure # they match the session whose mid was specified assert self.product_id == session.user.product_id, "product_id mismatch" - assert ( - self.product_user_id == session.user.product_user_id - ), "product_user_id mismatch" + assert self.product_user_id == session.user.product_user_id, ( + "product_user_id mismatch" + ) # validate the Session's mid is "recent" - assert (datetime.now(tz=UTC) - session.started) < timedelta( - minutes=90 - ), "expired session" + assert (datetime.now(tz=UTC) - session.started) < timedelta(minutes=90), ( + "expired session" + ) def model_dump_sql(self, **kwargs) -> dict[str, Any]: d = {} diff --git a/generalresearch/incite/mergers/foundations/enriched_wall.py b/generalresearch/incite/mergers/foundations/enriched_wall.py index 396c2be..70139c2 100644 --- a/generalresearch/incite/mergers/foundations/enriched_wall.py +++ b/generalresearch/incite/mergers/foundations/enriched_wall.py @@ -40,7 +40,6 @@ class EnrichedWallMergeItem(MergeCollectionItem): session_coll: SessionDFCollection, pg_config: PostgresConfig, client: Client | None = None, - client_resources: dict[str, Any] | None = None, ) -> None: ir: pd.Interval = self.interval @@ -160,7 +159,6 @@ class EnrichedWallMergeItem(MergeCollectionItem): ddf=ddf, is_partial=True, validate_after=False, - client_resources=client_resources, ) else: df = self.validate_df(df=df) @@ -169,7 +167,6 @@ class EnrichedWallMergeItem(MergeCollectionItem): client, ddf=ddf, is_partial=False, - client_resources=client_resources, ) diff --git a/pyproject.toml b/pyproject.toml index bb23838..13fa584 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -9,12 +9,12 @@ description = "Python Utilities for General Research" readme = "README.md" requires-python = ">=3.8" dependencies = [ - "fastapi", "Faker", "PyMySQL", "psycopg", "cachetools", "decorator", + "influxdb", "limits", "more-itertools", "numpy", diff --git a/test_utils/managers/gr/conftest.py b/test_utils/managers/gr/conftest.py index cc1053c..09e08f5 100644 --- a/test_utils/managers/gr/conftest.py +++ b/test_utils/managers/gr/conftest.py @@ -29,6 +29,8 @@ if TYPE_CHECKING: @pytest.fixture(scope="session") def gr_redis_config_db() -> str: + # need to update 'databases' in /etc/redis/redis.conf + # or this won't work and you'll have no indication why ... return str(randint(99, 1_023)) diff --git a/test_utils/managers/thl/conftest.py b/test_utils/managers/thl/conftest.py index 8ca4383..98dd574 100644 --- a/test_utils/managers/thl/conftest.py +++ b/test_utils/managers/thl/conftest.py @@ -92,6 +92,11 @@ def thl_redis_config( r.flushdb() +@pytest.fixture(scope="session") +def thl_redis_client(thl_redis_config): + return thl_redis_config.create_redis_client() + + @pytest.fixture(scope="session") def thl_web_rr(django_db_factory: Callable[..., PostgresDsn]) -> PostgresConfig: _dsn = django_db_factory("generalresearch.thl_django") diff --git a/test_utils/models/gr/conftest.py b/test_utils/models/gr/conftest.py index 1dbea0c..a48656b 100644 --- a/test_utils/models/gr/conftest.py +++ b/test_utils/models/gr/conftest.py @@ -12,7 +12,7 @@ from pydantic_extra_types.phone_numbers import PhoneNumber from generalresearch.models.custom_types import UUIDStr if TYPE_CHECKING: - from generalresearch.managers.gr.authentication import GRTokenManager, GRUserManager + from generalresearch.managers.gr.authentication import GRUserManager from generalresearch.managers.gr.business import ( BusinessAddressManager, BusinessBankAccountManager, @@ -289,9 +289,9 @@ def gr_user_token_factory( gr_user.prefetch_token(pg_config=gr_db) res = gr_user.token - assert ( - res is not None - ), "GRToken should exist after creation and prefetching" + assert res is not None, ( + "GRToken should exist after creation and prefetching" + ) return res else: @@ -335,8 +335,10 @@ def gr_membership_factory( @pytest.fixture() -def gr_membership(gr_membership_factory: Callable[..., Membership]) -> Membership: - return gr_membership_factory(save=True) +def gr_membership( + gr_membership_factory: Callable[..., Membership], gr_team: Team, gr_user: GRUser +) -> Membership: + return gr_membership_factory(gr_team=gr_team, gr_user=gr_user, save=True) @pytest.fixture() diff --git a/tests/conftest.py b/tests/conftest.py index 4777e15..b69d7ea 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -12,7 +12,6 @@ pytest_plugins = [ "test_utils.managers.contest.conftest", "test_utils.managers.gr.conftest", "test_utils.managers.ledger.conftest", - "test_utils.managers.network.conftest", "test_utils.managers.thl.conftest", "test_utils.managers.upk.conftest", # -- Models @@ -20,7 +19,6 @@ pytest_plugins = [ "test_utils.models.contest.conftest", "test_utils.models.gr.conftest", "test_utils.models.ledger.conftest", - "test_utils.models.network.conftest", "test_utils.models.thl.conftest", "test_utils.models.upk.conftest", # -- Marketplaces diff --git a/tests/models/custom_types/test_dsn.py b/tests/models/custom_types/test_dsn.py index 2aae579..eff02d3 100644 --- a/tests/models/custom_types/test_dsn.py +++ b/tests/models/custom_types/test_dsn.py @@ -1,14 +1,12 @@ 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 -if TYPE_CHECKING: - from generalresearch.models.custom_types import DaskDsn, SentryDsn +from generalresearch.models.custom_types import DaskDsn, SentryDsn # --- Test Pydantic Models --- @@ -23,7 +21,6 @@ class SettingsModel(BaseModel): class TestDaskDsn: - def test_base(self): from dask.distributed import Client diff --git a/tests/models/gr/test_authentication.py b/tests/models/gr/test_authentication.py index 059a0a4..881571c 100644 --- a/tests/models/gr/test_authentication.py +++ b/tests/models/gr/test_authentication.py @@ -10,7 +10,6 @@ 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.team import Team @@ -26,7 +25,6 @@ SSO_ISSUER = "" class TestGRUser: - def test_init(self, gr_user: GRUser): assert isinstance(gr_user, GRUser) @@ -43,7 +41,7 @@ class TestGRUser: def test_teams( self, gr_user: GRUser, - membership: Membership, + gr_membership: Membership, gr_db: PostgresConfig, gr_redis_config: RedisConfig, ): @@ -60,16 +58,16 @@ class TestGRUser: self, gr_user_token: GRToken, gr_user: GRUser, - membership: Membership, + gr_membership: Membership, product_factory: Callable[..., Product], - membership_factory: Callable[..., Membership], - team: Team, + gr_membership_factory: Callable[..., Membership], + gr_team: Team, thl_web_rr: PostgresConfig, gr_redis_config: RedisConfig, gr_db: PostgresConfig, ): - product_factory(team=team) - membership_factory(team=team, gr_user=gr_user) + product_factory(team=gr_team) + gr_membership_factory(team=gr_team, gr_user=gr_user) gr_user.prefetch_teams( pg_config=gr_db, @@ -82,8 +80,8 @@ class TestGRUser: self, gr_user: GRUser, product_factory: Callable[..., Product], - team: Team, - membership: Membership, + gr_team: Team, + gr_membership: Membership, gr_db: PostgresConfig, thl_web_rr: PostgresConfig, gr_redis_config: RedisConfig, @@ -94,15 +92,15 @@ class TestGRUser: # Create a new Team membership, and then create a Product that # is part of that team - membership.prefetch_team(pg_config=gr_db, redis_config=gr_redis_config) - assert isinstance(membership.team, Team) + gr_membership.prefetch_team(pg_config=gr_db, redis_config=gr_redis_config) + assert isinstance(gr_membership.team, Team) - p: Product = product_factory(team=team) + p: Product = product_factory(team=gr_team) assert p.id_int - assert team.uuid == membership.team.uuid - assert p.team_id == team.uuid - assert p.team_uuid == membership.team.uuid - assert gr_user.id == membership.user_id + assert gr_team.uuid == gr_membership.team.uuid + assert p.team_id == gr_team.uuid + assert p.team_uuid == gr_membership.team.uuid + assert gr_user.id == gr_membership.user_id gr_user.prefetch_products( pg_config=gr_db, @@ -115,7 +113,6 @@ class TestGRUser: class TestGRUserMethods: - def test_cache_key(self, gr_user: GRUser): assert isinstance(gr_user.cache_key, str) assert ":" in gr_user.cache_key @@ -124,13 +121,13 @@ class TestGRUserMethods: def test_to_redis( self, gr_user: GRUser, - team: Team, + gr_team: Team, gr_business: Business, product_factory: Callable[..., Product], - membership_factory: Callable[..., Membership], + gr_membership_factory: Callable[..., Membership], ): - product_factory(team=team, business=gr_business) - membership_factory(team=team, gr_user=gr_user) + product_factory(team=gr_team, business=gr_business) + gr_membership_factory(team=gr_team, gr_user=gr_user) res = gr_user.to_redis() assert isinstance(res, str) @@ -171,16 +168,16 @@ class TestGRUserMethods: gr_db: PostgresConfig, thl_web_rr: PostgresConfig, product_factory: Callable[..., Product], - team: Team, - membership_factory: Callable[..., Membership], + gr_team: Team, + gr_membership_factory: Callable[..., Membership], thl_redis_config: RedisConfig, ): 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) + p1 = product_factory(team=gr_team) + gr_membership_factory(team=gr_team, gr_user=gr_user) gr_user.set_cache( pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config @@ -206,10 +203,10 @@ class TestGRUserMethods: gr_db: PostgresConfig, thl_web_rr: PostgresConfig, product_factory: Callable[..., Product], - team: Team, + gr_team: Team, gr_redis_config: RedisConfig, ): - product_factory(team=team) + product_factory(team=gr_team) client = gr_redis_config.create_redis_client() gr_user.set_cache( @@ -227,10 +224,10 @@ class TestGRUserMethods: thl_web_rr: PostgresConfig, product_factory: Callable[..., Product], gr_business: Business, - team: Team, + gr_team: Team, gr_redis_config: RedisConfig, ): - product_factory(team=team, business=gr_business) + product_factory(team=gr_team, business=gr_business) gr_user.set_cache( pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config @@ -247,10 +244,10 @@ class TestGRUserMethods: gr_db: PostgresConfig, thl_web_rr: PostgresConfig, product_factory: Callable[..., Product], - team: Team, + gr_team: Team, gr_redis_config: RedisConfig, ): - product_factory(team=team) + product_factory(team=gr_team) gr_user.set_cache( pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config @@ -262,7 +259,6 @@ class TestGRUserMethods: class TestGRToken: - @pytest.fixture def gr_token(self, gr_user: GRUser): now = datetime.now(tz=UTC) @@ -290,7 +286,6 @@ class TestGRToken: class TestClaims: - def test_init(self): d = { diff --git a/tests/models/gr/test_business.py b/tests/models/gr/test_business.py index 030a214..5d0de4f 100644 --- a/tests/models/gr/test_business.py +++ b/tests/models/gr/test_business.py @@ -119,11 +119,11 @@ class TestBusiness: def duration(self) -> timedelta | None: return None - def test_init(self, business: Business): + def test_init(self, gr_business: Business): - assert isinstance(business, Business) - assert isinstance(business.id, int) - assert isinstance(business.uuid, str) + assert isinstance(gr_business, Business) + assert isinstance(gr_business.id, int) + assert isinstance(gr_business.uuid, str) def test_str_and_repr( self, @@ -208,17 +208,17 @@ class TestBusiness: def test_addresses( self, - business: Business, + gr_business: Business, gr_db: PostgresConfig, ): from generalresearch.models.gr.business import BusinessAddress - assert business.addresses is None + assert gr_business.addresses is None - business.prefetch_addresses(pg_config=gr_db) - assert isinstance(business.addresses, list) - assert len(business.addresses) == 1 - assert isinstance(business.addresses[0], BusinessAddress) + gr_business.prefetch_addresses(pg_config=gr_db) + assert isinstance(gr_business.addresses, list) + assert len(gr_business.addresses) == 1 + assert isinstance(gr_business.addresses[0], BusinessAddress) def test_teams( self, @@ -674,8 +674,6 @@ class TestBusinessBalance: started=start + timedelta(days=2), ) - payout_event_manager.set_account_lookup_table(thl_lm=thl_ledger_manager) - brokerage_product_payout_event_factory( product=u1.product, amount=USDCent(5), @@ -770,7 +768,6 @@ class TestBusinessBalance: wall_req_cpi=Decimal("2.50"), started=start + timedelta(days=2), ) - payout_event_manager.set_account_lookup_table(thl_lm=thl_ledger_manager) brokerage_product_payout_event_factory( product=u1.product, @@ -887,7 +884,6 @@ class TestBusinessBalance: wall_req_cpi=Decimal(".75"), started=start + timedelta(days=1), ) - payout_event_manager.set_account_lookup_table(thl_lm=thl_ledger_manager) brokerage_product_payout_event_factory( product=u1.product, amount=USDCent(71), @@ -1041,7 +1037,6 @@ class TestBusinessBalance: wall_req_cpi=Decimal("2.50"), started=start + timedelta(days=2), ) - payout_event_manager.set_account_lookup_table(thl_lm=thl_ledger_manager) brokerage_product_payout_event_factory( product=u1.product, diff --git a/tests/models/gr/test_team.py b/tests/models/gr/test_team.py index 8ebedb6..b5f1781 100644 --- a/tests/models/gr/test_team.py +++ b/tests/models/gr/test_team.py @@ -42,7 +42,6 @@ if TYPE_CHECKING: class TestTeam: - def test_init(self, gr_team: Team): assert isinstance(gr_team, Team) @@ -54,7 +53,7 @@ class TestTeam: ): assert gr_team.memberships is None - gr_team.prefetch_memberships(membership_manager=gr_membership_manager) + gr_team.prefetch_memberships(gr_membership_manager=gr_membership_manager) assert isinstance(gr_team.memberships, list) assert len(gr_team.memberships) == 0 @@ -67,7 +66,7 @@ class TestTeam: ): assert gr_team.memberships is None - gr_team.prefetch_memberships(membership_manager=gr_membership_manager) + gr_team.prefetch_memberships(gr_membership_manager=gr_membership_manager) assert isinstance(gr_team.memberships, list) assert len(gr_team.memberships) == 1 assert gr_team.memberships[0].user_id == gr_user.id @@ -75,7 +74,7 @@ class TestTeam: # Create another new Membership gr_membership_manager.create(team=gr_team, gr_user=gr_user_factory()) assert len(gr_team.memberships) == 1 - gr_team.prefetch_memberships(membership_manager=gr_membership_manager) + gr_team.prefetch_memberships(gr_membership_manager=gr_membership_manager) assert len(gr_team.memberships) == 2 def test_gr_users( @@ -146,7 +145,6 @@ class TestTeam: class TestTeamMethods: - def test_cache_key(self, gr_team: Team): assert isinstance(gr_team.cache_key, str) assert ":" in gr_team.cache_key diff --git a/tests/models/test_finance.py b/tests/models/test_finance.py index 502c596..a1da961 100644 --- a/tests/models/test_finance.py +++ b/tests/models/test_finance.py @@ -13,9 +13,6 @@ import pytest from dask.distributed import Client as DaskClient # noinspection PyUnresolvedReferences -from distributed.utils_test import ( - client_no_amm, -) from faker import Faker from generalresearch.incite.schemas.mergers.pop_ledger import ( @@ -26,8 +23,6 @@ from generalresearch.models.thl.finance import ( POPFinancial, ProductBalances, ) -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 @@ -43,7 +38,6 @@ fake = Faker() class TestProductBalanceInitialize: - def test_unknown_fields(self): with pytest.raises(expected_exception=ValueError): ProductBalances.model_validate( @@ -251,7 +245,6 @@ class TestProductBalanceInitialize: class TestBusinessBalanceInitialize: - def test_validate_product_ids(self): instance1 = ProductBalances.model_validate( {"bp_payment.CREDIT": 500, "bp_adjustment.DEBIT": 40} @@ -668,9 +661,11 @@ class TestBusinessBalanceInitialize: ), ) class TestProductFinanceData: - def test_base( self, + ledger_collection: LedgerDFCollection, + pop_ledger_merge, + client_no_amm, duration: timedelta, product: Product, user_factory: Callable[..., User], @@ -681,9 +676,9 @@ class TestProductFinanceData: # -- Build & Setup u: User = user_factory(product=product, created=ledger_collection.start) + assert u.product for item in ledger_collection.items: - for _ in range(3): rand_item_time = fake.date_time_between( start_date=item.start, @@ -737,7 +732,6 @@ class TestProductFinanceData: class TestPOPFinancialData: - def test_base( self, client_no_amm: DaskClient, diff --git a/tests/models/thl/test_contest/test_leaderboard_contest.py b/tests/models/thl/test_contest/test_leaderboard_contest.py index c49776b..f787bdf 100644 --- a/tests/models/thl/test_contest/test_leaderboard_contest.py +++ b/tests/models/thl/test_contest/test_leaderboard_contest.py @@ -33,7 +33,7 @@ class TestLeaderboardContest(TestContest): @pytest.fixture def leaderboard_contest( - self, product: Product, thl_redis: Redis, user_manager: UserManager + self, product: Product, thl_redis_client: Redis, user_manager: UserManager ) -> LeaderboardContest: board_key = f"leaderboard:{product.uuid}:us:weekly:2025-05-26:complete_count" @@ -67,14 +67,14 @@ class TestLeaderboardContest(TestContest): ), ], ) - c._redis_client = thl_redis + c._redis_client = thl_redis_client c._user_manager = user_manager return c def test_init( self, leaderboard_contest: LeaderboardContest, - thl_redis: Redis, + thl_redis_client: Redis, user_1: User, user_2: User, ): @@ -82,7 +82,7 @@ class TestLeaderboardContest(TestContest): assert leaderboard_contest.end_condition.ends_at is not None lbm = LeaderboardManager( - redis_client=thl_redis, + redis_client=thl_redis_client, board_code=model.board_code, country_iso=model.country_iso, freq=model.freq, diff --git a/tests/models/thl/test_payout_format.py b/tests/models/thl/test_payout_format.py index 56eafe3..fe7aea5 100644 --- a/tests/models/thl/test_payout_format.py +++ b/tests/models/thl/test_payout_format.py @@ -1,20 +1,14 @@ 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 223430f..25affcf 100644 --- a/tests/models/thl/test_product.py +++ b/tests/models/thl/test_product.py @@ -48,7 +48,6 @@ if TYPE_CHECKING: class TestProduct: - def test_init(self): # By default, just a Pydantic instance doesn't have an id_int instance = Product.model_validate( @@ -70,13 +69,13 @@ class TestProduct: # By default, just a Pydantic instance doesn't have an id_int instance = product_factory() assert isinstance(instance.id_int, int) + assert isinstance(instance, Product) res = instance.model_dump_json() - assert isinstance(res, Product) # we json skip & exclude - res = instance.model_dump() - assert isinstance(res, Product) + p = Product.model_validate_json(res) + assert isinstance(p, Product) def test_redirect_url(self): p = Product.model_validate( @@ -150,12 +149,6 @@ class TestProduct: redirect_url="https://www.google.com/hey", ) - assert isinstance(p.payout_config.payout_transformation, PayoutTransformation) - assert isinstance( - p.payout_config.payout_transformation.kwargs, - PayoutTransformationPercentArgs, - ) - p.payout_config.payout_transformation = PayoutTransformation.model_validate( { "f": "payout_transformation_percent", @@ -598,7 +591,6 @@ class TestGlobalProductConfigFor: class TestProductFinancials: - @pytest.fixture def start(self) -> datetime: return datetime(year=2018, month=3, day=14, hour=0, tzinfo=UTC) @@ -639,7 +631,6 @@ class TestProductFinancials: u1: User = user_factory(product=p1) 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( @@ -818,7 +809,6 @@ class TestProductFinancials: class TestProductBalance: - @pytest.fixture def start(self) -> datetime: return datetime(year=2018, month=3, day=14, hour=0, tzinfo=UTC) @@ -867,7 +857,6 @@ 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_ledger_manager) brokerage_product_payout_event_factory( product=product, amount=USDCent(71), @@ -928,7 +917,6 @@ 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) brokerage_product_payout_event_factory( product=product, amount=USDCent(71), @@ -947,7 +935,6 @@ class TestProductBalance: class TestProductPOPFinancial: - @pytest.fixture def start(self) -> datetime: return datetime(year=2018, month=3, day=14, hour=0, tzinfo=UTC) @@ -1020,7 +1007,6 @@ class TestProductPOPFinancial: class TestProductCache: - @pytest.fixture def start(self) -> datetime: return datetime(year=2018, month=3, day=14, hour=0, tzinfo=UTC) @@ -1143,7 +1129,6 @@ class TestProductCache: ) # 2. Payout - payout_event_manager.set_account_lookup_table(thl_lm=thl_ledger_manager) brokerage_product_payout_event_factory( product=product, amount=USDCent(71), -- cgit v1.2.3 From c720350aaf92d68d2f84ba05702b05acb448fa00 Mon Sep 17 00:00:00 2001 From: stuppie Date: Fri, 4 Sep 2026 13:03:31 -0600 Subject: fix gr_business_bank_account fixture. fix some more tests --- test_utils/models/gr/conftest.py | 23 ++++++++++------------- tests/models/gr/test_authentication.py | 6 +++--- 2 files changed, 13 insertions(+), 16 deletions(-) (limited to 'test_utils/models') diff --git a/test_utils/models/gr/conftest.py b/test_utils/models/gr/conftest.py index a48656b..859aaa4 100644 --- a/test_utils/models/gr/conftest.py +++ b/test_utils/models/gr/conftest.py @@ -10,9 +10,10 @@ from pydantic import PositiveInt from pydantic_extra_types.phone_numbers import PhoneNumber from generalresearch.models.custom_types import UUIDStr +from generalresearch.models.gr.definitions import TransferMethod if TYPE_CHECKING: - from generalresearch.managers.gr.authentication import GRUserManager + from generalresearch.managers.gr.authentication import GRTokenManager, GRUserManager from generalresearch.managers.gr.business import ( BusinessAddressManager, BusinessBankAccountManager, @@ -25,7 +26,6 @@ if TYPE_CHECKING: BusinessAddress, BusinessBankAccount, ) - 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 @@ -74,15 +74,11 @@ def gr_business_bank_account_factory( @pytest.fixture -def gr_business_bank_account(gr_business_factory: Callable[..., Business]) -> Business: - return gr_business_factory(save=True) - - -@pytest.fixture -def unsaved_gr_business_bank_account( - gr_business_factory: Callable[..., Business], -) -> Business: - return gr_business_factory(save=False) +def gr_business_bank_account( + gr_business_bank_account_factory: Callable[..., BusinessBankAccount], + gr_business: Business, +) -> BusinessBankAccount: + return gr_business_bank_account_factory(save=True, business_id=gr_business.id) # --- Business Address --- @@ -277,7 +273,7 @@ def unsaved_gr_user( @pytest.fixture def gr_user_token_factory( - gr_user: GRUser, gr_user_token_manager: GRUser, gr_db: PostgresConfig + gr_user: GRUser, gr_token_manager: GRTokenManager, gr_db: PostgresConfig ) -> Callable[..., GRToken]: def _inner( @@ -285,7 +281,8 @@ def gr_user_token_factory( ) -> GRToken: if save: - gr_user_token_manager.create(user_id=gr_user.id) + assert gr_user.id + gr_token_manager.create(user_id=gr_user.id) gr_user.prefetch_token(pg_config=gr_db) res = gr_user.token diff --git a/tests/models/gr/test_authentication.py b/tests/models/gr/test_authentication.py index 881571c..7ff44d0 100644 --- a/tests/models/gr/test_authentication.py +++ b/tests/models/gr/test_authentication.py @@ -67,7 +67,7 @@ class TestGRUser: gr_db: PostgresConfig, ): product_factory(team=gr_team) - gr_membership_factory(team=gr_team, gr_user=gr_user) + gr_membership_factory(gr_team=gr_team, gr_user=gr_user) gr_user.prefetch_teams( pg_config=gr_db, @@ -127,7 +127,7 @@ class TestGRUserMethods: gr_membership_factory: Callable[..., Membership], ): product_factory(team=gr_team, business=gr_business) - gr_membership_factory(team=gr_team, gr_user=gr_user) + gr_membership_factory(gr_team=gr_team, gr_user=gr_user) res = gr_user.to_redis() assert isinstance(res, str) @@ -177,7 +177,7 @@ class TestGRUserMethods: client = gr_redis_config.create_redis_client() p1 = product_factory(team=gr_team) - gr_membership_factory(team=gr_team, gr_user=gr_user) + gr_membership_factory(gr_team=gr_team, gr_user=gr_user) gr_user.set_cache( pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config -- cgit v1.2.3 From 4ac65aa2a3ffffcd0746588e3073548ce39aa245 Mon Sep 17 00:00:00 2001 From: stuppie Date: Fri, 4 Sep 2026 13:15:06 -0600 Subject: user_factory silently ignored passed in product_id and generates a new product --- test_utils/models/thl/conftest.py | 2 ++ 1 file changed, 2 insertions(+) (limited to 'test_utils/models') diff --git a/test_utils/models/thl/conftest.py b/test_utils/models/thl/conftest.py index 9cb68be..ff3bebd 100644 --- a/test_utils/models/thl/conftest.py +++ b/test_utils/models/thl/conftest.py @@ -646,6 +646,8 @@ def user_factory( ) -> User: if save: if product is None: + if product_id: + raise ValueError("this is broken") product = product_factory() product_user_id = product_user_id or uuid4().hex -- cgit v1.2.3 From 63f9d47b2774993e5e749d9d9c64915ad0d2de50 Mon Sep 17 00:00:00 2001 From: stuppie Date: Fri, 4 Sep 2026 14:22:16 -0600 Subject: session fixture to use the User fixture otherwise tests that use session and user wont match. ip_geoname_factory default save. ip_information_factory missing fixture --- test_utils/models/thl/conftest.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) (limited to 'test_utils/models') diff --git a/test_utils/models/thl/conftest.py b/test_utils/models/thl/conftest.py index ff3bebd..8b3ab32 100644 --- a/test_utils/models/thl/conftest.py +++ b/test_utils/models/thl/conftest.py @@ -239,9 +239,9 @@ def bare_session_factory( @pytest.fixture() -def bare_session(bare_session_factory: Callable[..., Session]) -> Session: +def bare_session(bare_session_factory: Callable[..., Session], user) -> Session: # A session with no wall events - return bare_session_factory() + return bare_session_factory(user=user) @pytest.fixture @@ -447,7 +447,7 @@ def ip_geoname_factory( ) -> Callable[..., IPGeoname]: def _inner( - save: bool, + save: bool = True, geoname_id: PositiveInt | None = None, continent_code: str | None = None, continent_name: str | None = None, @@ -496,7 +496,7 @@ def unsaved_ip_geoname(ip_geoname_factory: Callable[..., IPGeoname]) -> IPGeonam # --- IP Information --- - +@pytest.fixture def ip_information_factory( ipinformation_manager: IPInformationManager, ) -> Callable[..., IPInformation]: -- cgit v1.2.3 From 0645e930703a6fa01ef913a3a93720aedbe6a9f5 Mon Sep 17 00:00:00 2001 From: stuppie Date: Sun, 6 Sep 2026 19:32:31 -0600 Subject: GrlIqDataManager: something changed but I can't really figure out what; not sure if its a test issue or really broken, but either way results and category_result should be excluded from model dump --- generalresearch/grliq/managers/forensic_data.py | 55 ++++++++++++++----------- test_utils/grliq/conftest.py | 2 +- test_utils/models/thl/conftest.py | 4 +- 3 files changed, 33 insertions(+), 28 deletions(-) (limited to 'test_utils/models') diff --git a/generalresearch/grliq/managers/forensic_data.py b/generalresearch/grliq/managers/forensic_data.py index f523b6e..0810723 100644 --- a/generalresearch/grliq/managers/forensic_data.py +++ b/generalresearch/grliq/managers/forensic_data.py @@ -17,13 +17,11 @@ from generalresearch.grliq.models.forensic_result import ( from generalresearch.models.custom_types import UUIDStr if TYPE_CHECKING: - from generalresearch.models.thl.user import User from generalresearch.pg_helper import PostgresConfig class GrlIqDataManager: - def __init__(self, postgres_config: PostgresConfig): self.postgres_config = postgres_config @@ -36,7 +34,15 @@ class GrlIqDataManager: is_attempt_allowed: bool | None = None, ) -> GrlIqData: - data = iq_data.model_dump_sql(exclude={"events", "mouse_events", "timing_data"}) + data = iq_data.model_dump_sql( + exclude={ + "events", + "mouse_events", + "timing_data", + "results", + "category_result", + } + ) data["result_data"] = None if result_data: @@ -476,16 +482,16 @@ class GrlIqDataManager: product_ids = None if product_ids: - assert ( - users is None and user is None and product_id is None - ), "user, users, product_id, and product_ids are mutually exclusive" + assert users is None and user is None and product_id is None, ( + "user, users, product_id, and product_ids are mutually exclusive" + ) params["product_ids"] = list(set(product_ids)) filters.append("d.product_id = ANY(%(product_ids)s::UUID[])") if product_id: - assert ( - users is None and user is None and product_ids is None - ), "user, users, product_id, and product_ids are mutually exclusive" + assert users is None and user is None and product_ids is None, ( + "user, users, product_id, and product_ids are mutually exclusive" + ) params["product_id"] = product_id filters.append("d.product_id = %(product_id)s") @@ -506,12 +512,12 @@ class GrlIqDataManager: ) if created_between: - assert ( - created_after is None - ), "Cannot pass both created_after and created_between" - assert ( - created_before is None - ), "Cannot pass both created_before and created_between" + assert created_after is None, ( + "Cannot pass both created_after and created_between" + ) + assert created_before is None, ( + "Cannot pass both created_before and created_between" + ) params["created_after"] = created_between[0] params["created_before"] = created_between[1] filters.append( @@ -519,9 +525,9 @@ class GrlIqDataManager: ) if user: - assert ( - product_ids is None and users is None - ), "user, users, and product_ids are mutually exclusive" + assert product_ids is None and users is None, ( + "user, users, and product_ids are mutually exclusive" + ) params["product_id"] = user.product_id params["product_user_id"] = user.product_user_id filters.append( @@ -529,9 +535,9 @@ class GrlIqDataManager: ) if users: - assert ( - product_ids is None and user is None - ), "user, users, and product_ids are mutually exclusive" + assert product_ids is None and user is None, ( + "user, users, and product_ids are mutually exclusive" + ) user_args = ", ".join( [f"(%(bp_{i})s, %(bpuid_{i})s)" for i in range(len(users))] ) @@ -649,9 +655,9 @@ class GrlIqDataManager: if product_ids: # It doesn't use the (product_id, created_at) index with multiple product_ids - assert ( - offset == 0 - ), "Cannot paginate using product_ids, use product_id instead" + assert offset == 0, ( + "Cannot paginate using product_ids, use product_id instead" + ) filter_str, params = self.make_filter_str( session_uuid=session_uuid, @@ -682,7 +688,6 @@ class GrlIqDataManager: res: list[dict[str, Any]] = c.fetchall() # type: ignore for x in res: - if "data" in x: self.temporary_add_missing_fields(x["data"]) x["data"]["id"] = x["id"] diff --git a/test_utils/grliq/conftest.py b/test_utils/grliq/conftest.py index f399a99..249b068 100644 --- a/test_utils/grliq/conftest.py +++ b/test_utils/grliq/conftest.py @@ -90,7 +90,7 @@ def grliq_data_factory( """ if save: - res: GrlIqData = grliq_data_list[int(is_attempt_allowed)]["data"] + res: dict = grliq_data_list[int(is_attempt_allowed)] product_id = product_id or uuid4().hex product_user_id = product_user_id or uuid4().hex diff --git a/test_utils/models/thl/conftest.py b/test_utils/models/thl/conftest.py index 8b3ab32..fe2b4f5 100644 --- a/test_utils/models/thl/conftest.py +++ b/test_utils/models/thl/conftest.py @@ -498,7 +498,7 @@ def unsaved_ip_geoname(ip_geoname_factory: Callable[..., IPGeoname]) -> IPGeonam @pytest.fixture def ip_information_factory( - ipinformation_manager: IPInformationManager, + ip_information_manager: IPInformationManager, ) -> Callable[..., IPInformation]: def _inner( @@ -530,7 +530,7 @@ def ip_information_factory( ) -> IPInformation: if save: - return ipinformation_manager.create( + return ip_information_manager.create( ip=ip or fake.ipv4_public(), geoname_id=geoname_id, country_iso=country_iso or fake.country_code(), -- cgit v1.2.3 From 9002e4ea94ce790210ce42f4950a9e63ee8ca27b Mon Sep 17 00:00:00 2001 From: stuppie Date: Mon, 7 Sep 2026 10:43:20 -0600 Subject: fix more tests --- test_utils/managers/thl/conftest.py | 2 +- test_utils/models/thl/conftest.py | 10 +++------- tests/managers/thl/test_task_status.py | 6 +++--- tests/managers/thl/test_user_manager/test_base.py | 2 +- tests/managers/thl/test_user_manager/test_redis.py | 2 +- tests/managers/thl/test_user_streak.py | 12 ++++-------- tests/managers/thl/test_userhealth.py | 8 ++++---- 7 files changed, 17 insertions(+), 25 deletions(-) (limited to 'test_utils/models') diff --git a/test_utils/managers/thl/conftest.py b/test_utils/managers/thl/conftest.py index 98dd574..355a39d 100644 --- a/test_utils/managers/thl/conftest.py +++ b/test_utils/managers/thl/conftest.py @@ -240,7 +240,7 @@ def mysql_user_manager(thl_web_rw: PostgresConfig) -> MysqlUserManager: @pytest.fixture(scope="session") def redis_user_manager(thl_redis_config: RedisConfig) -> RedisUserManager: - return RedisUserManager(redis_dsn=thl_redis_config) + return RedisUserManager(redis_dsn=thl_redis_config.dsn) @pytest.fixture(scope="session") diff --git a/test_utils/models/thl/conftest.py b/test_utils/models/thl/conftest.py index fe2b4f5..376891c 100644 --- a/test_utils/models/thl/conftest.py +++ b/test_utils/models/thl/conftest.py @@ -496,6 +496,7 @@ def unsaved_ip_geoname(ip_geoname_factory: Callable[..., IPGeoname]) -> IPGeonam # --- IP Information --- + @pytest.fixture def ip_information_factory( ip_information_manager: IPInformationManager, @@ -833,13 +834,8 @@ def audit_log_factory(audit_log_manager: AuditLogManager) -> Callable[..., Audit @pytest.fixture() -def audit_log(auditlog_factory: Callable[..., AuditLog]) -> AuditLog: - return auditlog_factory(save=True) - - -@pytest.fixture() -def unsaved_audit_log(auditlog_factory: Callable[..., AuditLog]) -> AuditLog: - return auditlog_factory(save=False) +def audit_log(audit_log_factory: Callable[..., AuditLog], user: User) -> AuditLog: + return audit_log_factory(user_id=user.user_id) # --- --- diff --git a/tests/managers/thl/test_task_status.py b/tests/managers/thl/test_task_status.py index 4a401fa..11edc99 100644 --- a/tests/managers/thl/test_task_status.py +++ b/tests/managers/thl/test_task_status.py @@ -39,7 +39,7 @@ start3 = datetime(2023, 2, 3, tzinfo=UTC) finish3 = start3 + timedelta(minutes=5) -@pytest.fixture(scope="session") +@pytest.fixture() def bp1( product_factory: Callable[..., Product], product_manager: ProductManager ) -> Product: @@ -50,7 +50,7 @@ def bp1( ) -@pytest.fixture(scope="session") +@pytest.fixture() def bp2( product_factory: Callable[..., Product], product_manager: ProductManager ) -> Product: @@ -66,7 +66,7 @@ def bp2( ) -@pytest.fixture(scope="session") +@pytest.fixture() def bp3( product_factory: Callable[..., Product], product_manager: ProductManager ) -> Product: diff --git a/tests/managers/thl/test_user_manager/test_base.py b/tests/managers/thl/test_user_manager/test_base.py index c69f297..5d12052 100644 --- a/tests/managers/thl/test_user_manager/test_base.py +++ b/tests/managers/thl/test_user_manager/test_base.py @@ -305,7 +305,7 @@ class TestUserManagerMethods: assert len(res) == 0 msg = uuid4().hex - user_manager.audit_log(user=user, level=30, event_type=msg) + user_manager.audit_log(audit_log_manager, user=user, level=30, event_type=msg) res = audit_log_manager.filter_by_user_id(user_id=user.user_id) assert len(res) == 1 diff --git a/tests/managers/thl/test_user_manager/test_redis.py b/tests/managers/thl/test_user_manager/test_redis.py index e51aae9..f6b59c9 100644 --- a/tests/managers/thl/test_user_manager/test_redis.py +++ b/tests/managers/thl/test_user_manager/test_redis.py @@ -18,7 +18,7 @@ if TYPE_CHECKING: class TestUserManagerRedis: def test_get_notset(self, redis_user_manager: RedisUserManager, user: User): - redis_user_manager.clear_user_inmemory_cache(user=user) + redis_user_manager.clear_user(user=user) assert redis_user_manager.get_user(user_id=user.user_id) is None def test_get_user_id(self, redis_user_manager: RedisUserManager, user: User): diff --git a/tests/managers/thl/test_user_streak.py b/tests/managers/thl/test_user_streak.py index 59dee2d..d99b2b8 100644 --- a/tests/managers/thl/test_user_streak.py +++ b/tests/managers/thl/test_user_streak.py @@ -112,10 +112,8 @@ 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 = session_manager.create(started=start, country_iso="us", user=user) session_manager.finish_with_status( session, finished=start + timedelta(minutes=1), @@ -128,10 +126,8 @@ def create_session_complete( 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 = session_manager.create(started=start, country_iso="us", user=user) session_manager.finish_with_status( session, finished=start + timedelta(minutes=1), @@ -153,7 +149,7 @@ def test_user_streaks_active_broken( user: User, session_manager: SessionManager, broken_active_streak: list[UserStreak], - session_factory: Callable[..., Session], + bare_session_factory: Callable[..., Session], wall_factory: Callable[..., Wall], ): # Testing active streak, but broken (not today or yesterday) @@ -161,7 +157,7 @@ def test_user_streaks_active_broken( end1 = start1 + timedelta(minutes=1) # abandon counts as inactive - session = session_factory(started=start1, country_iso="us", user=user) + session = bare_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 a86361a..268b110 100644 --- a/tests/managers/thl/test_userhealth.py +++ b/tests/managers/thl/test_userhealth.py @@ -319,7 +319,7 @@ class TestUserIpHistoryManager: ip_geoname: IPGeoname, ): ip = fake.ipv4_public() - ip_information_factory(ip=ip, geoname=ip_geoname, is_anonymous=True) + ip_information_factory(ip=ip, geoname_id=ip_geoname.geoname_id, is_anonymous=True) ipr1: IPRecord = ip_record_factory(user_id=user.user_id, ip=ip) ipr = user_iphistory_manager.get_user_latest_ip_record(user=user) @@ -330,7 +330,7 @@ class TestUserIpHistoryManager: assert ipr.information.lookup_prefix == "/32" ip = fake.ipv6() - ip_information_factory(ip=ip, geoname=ip_geoname) + ip_information_factory(ip=ip, geoname_id=ip_geoname.geoname_id) ipr2: IPRecord = ip_record_factory(user_id=user.user_id, ip=ip) ipr = user_iphistory_manager.get_user_latest_ip_record(user=user) @@ -386,7 +386,7 @@ class TestUserIpHistoryManager: assert ipr.information is None assert not ipr.is_anonymous - ip_information_factory(ip=ip, geoname=ip_geoname, is_anonymous=True) + ip_information_factory(ip=ip, geoname_id=ip_geoname.geoname_id, is_anonymous=True) iph = user_iphistory_manager.get_user_ip_history(user_id=user.user_id) assert isinstance(iph, UserIPHistory) assert isinstance(iph.ips, list) @@ -414,7 +414,7 @@ class TestUserIpHistoryManager: assert ipr.information is None assert not ipr.is_anonymous - ip_information_factory(ip=ip, geoname=ip_geoname, is_anonymous=True) + ip_information_factory(ip=ip, geoname_id=ip_geoname.geoname_id, is_anonymous=True) iph = user_iphistory_manager.get_user_ip_history(user_id=user.user_id) assert isinstance(iph, UserIPHistory) assert isinstance(iph.ips, list) -- cgit v1.2.3 From 3338f74a94d0624bf894ebb35bd1bcfca268216e Mon Sep 17 00:00:00 2001 From: stuppie Date: Mon, 7 Sep 2026 11:12:04 -0600 Subject: fix more tests. Fix survey score optional field --- generalresearch/models/gr/team.py | 2 +- generalresearch/models/thl/survey/buyer.py | 3 +- test_utils/models/gr/conftest.py | 37 ++++------ test_utils/models/thl/conftest.py | 12 ++-- tests/models/gr/test_authentication.py | 2 + tests/models/gr/test_business.py | 80 +++++----------------- tests/models/gr/test_team.py | 13 ++-- .../thl/test_contest/test_leaderboard_contest.py | 4 +- tests/models/thl/test_product.py | 12 +--- 9 files changed, 52 insertions(+), 113 deletions(-) (limited to 'test_utils/models') diff --git a/generalresearch/models/gr/team.py b/generalresearch/models/gr/team.py index aaa5869..8d23bc5 100644 --- a/generalresearch/models/gr/team.py +++ b/generalresearch/models/gr/team.py @@ -273,7 +273,7 @@ class Team(BaseModel): self.prefetch_products(product_manager=product_manager) self.prefetch_gr_users(gr_user_manager=gr_user_manager) self.prefetch_businesses(gr_business_manager=gr_business_manager) - self.prefetch_memberships(membership_manager=gr_membership_manager) + self.prefetch_memberships(gr_membership_manager=gr_membership_manager) rc = redis_config.create_redis_client() mapping = self.model_dump(mode="json") diff --git a/generalresearch/models/thl/survey/buyer.py b/generalresearch/models/thl/survey/buyer.py index 6a67ed7..91102b4 100644 --- a/generalresearch/models/thl/survey/buyer.py +++ b/generalresearch/models/thl/survey/buyer.py @@ -177,9 +177,10 @@ class BuyerCountryStat(BaseModel): ) # ---- Scoring ---- - score: float = Field( + score: float | None = Field( description="Composite score calculated from all of the individual features", examples=[-5.329389837486194], + default=None, ) @model_validator(mode="after") diff --git a/test_utils/models/gr/conftest.py b/test_utils/models/gr/conftest.py index 859aaa4..f5dcaa1 100644 --- a/test_utils/models/gr/conftest.py +++ b/test_utils/models/gr/conftest.py @@ -91,7 +91,6 @@ def gr_business_address_factory( def _inner( business_id: PositiveInt, - save: bool = True, uuid: UUIDStr | None = None, line_1: str | None = None, line_2: str | None = None, @@ -110,36 +109,26 @@ def gr_business_address_factory( phone_number = None country = country or "US" - if save: - return gr_business_address_manager.create( - business_id=business_id, - uuid=uuid, - line_1=line_1, - line_2=line_2, - city=city, - state=state, - postal_code=postal_code, - phone_number=phone_number, - country=country, - ) - else: - raise ValueError("Unsaved BusinessAddress not supported yet") + return gr_business_address_manager.create( + business_id=business_id, + uuid=uuid, + line_1=line_1, + line_2=line_2, + city=city, + state=state, + postal_code=postal_code, + phone_number=phone_number, + country=country, + ) return _inner @pytest.fixture def gr_business_address( - gr_business_address_factory: Callable[..., BusinessAddress], -) -> BusinessAddress: - return gr_business_address_factory(save=True) - - -@pytest.fixture -def unsaved_gr_business_address( - gr_business_address_factory: Callable[..., BusinessAddress], + gr_business_address_factory: Callable[..., BusinessAddress], gr_business: Business ) -> BusinessAddress: - return gr_business_address_factory(save=False) + return gr_business_address_factory(business_id=gr_business.id) # --- Business --- diff --git a/test_utils/models/thl/conftest.py b/test_utils/models/thl/conftest.py index 376891c..e09eadd 100644 --- a/test_utils/models/thl/conftest.py +++ b/test_utils/models/thl/conftest.py @@ -14,7 +14,10 @@ from grip_client.enums import AccessType from pydantic import PositiveInt from generalresearch.currency import USDCent -from generalresearch.managers.thl.payout import UserPayoutEventManager +from generalresearch.managers.thl.payout import ( + BusinessPayoutEventManager, + UserPayoutEventManager, +) from generalresearch.models.custom_types import ( AwareDatetimeISO, IPvAnyAddressStr, @@ -38,9 +41,6 @@ if TYPE_CHECKING: IPInformationManager, ) from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager - from generalresearch.managers.thl.payout import ( - BrokerageProductPayoutEventManager, - ) from generalresearch.managers.thl.product import ProductManager from generalresearch.managers.thl.session import SessionManager from generalresearch.managers.thl.user_manager.user_manager import UserManager @@ -777,7 +777,7 @@ def unsaved_user_payout_event( @pytest.fixture def brokerage_product_payout_event_factory( thl_ledger_manager: ThlLedgerManager, - brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager, + business_payout_event_manager: BusinessPayoutEventManager, product_factory: Callable[..., Product], ) -> Callable[..., BrokerageProductPayoutEvent]: @@ -791,7 +791,7 @@ def brokerage_product_payout_event_factory( product = product or product_factory() amount = amount or USDCent(randint(1, 99_99)) - return brokerage_product_payout_event_manager.create_bp_payout_event( + return business_payout_event_manager.create_bp_payout_event( thl_ledger_manager=thl_ledger_manager, product=product, amount=amount, diff --git a/tests/models/gr/test_authentication.py b/tests/models/gr/test_authentication.py index 7ff44d0..21e07a4 100644 --- a/tests/models/gr/test_authentication.py +++ b/tests/models/gr/test_authentication.py @@ -205,6 +205,7 @@ class TestGRUserMethods: product_factory: Callable[..., Product], gr_team: Team, gr_redis_config: RedisConfig, + gr_membership, ): product_factory(team=gr_team) client = gr_redis_config.create_redis_client() @@ -246,6 +247,7 @@ class TestGRUserMethods: product_factory: Callable[..., Product], gr_team: Team, gr_redis_config: RedisConfig, + gr_membership, ): product_factory(team=gr_team) diff --git a/tests/models/gr/test_business.py b/tests/models/gr/test_business.py index 5d0de4f..e38850d 100644 --- a/tests/models/gr/test_business.py +++ b/tests/models/gr/test_business.py @@ -58,7 +58,6 @@ if TYPE_CHECKING: class TestBusinessBankAccount: - def test_init( self, gr_business: Business, @@ -93,13 +92,11 @@ class TestBusinessBankAccount: class TestBusinessAddress: - - def test_init(self, business_address: BusinessAddress): - assert isinstance(business_address, BusinessAddress) + def test_init(self, gr_business_address: BusinessAddress): + assert isinstance(gr_business_address, BusinessAddress) class TestBusinessContact: - def test_init(self): bc = BusinessContact(name="abc", email="test@abc.com") @@ -173,9 +170,6 @@ class TestBusiness: assert "Ledger Accounts: 2" in res3 # -- need some tx to make these interesting - business_payout_event_manager.set_account_lookup_table( - thl_lm=thl_ledger_manager - ) session_with_tx_factory( user=u1, wall_req_cpi=Decimal("2.50"), @@ -185,8 +179,6 @@ class TestBusiness: product=p1, amount=USDCent(50), created=start + timedelta(days=4), - skip_wallet_balance_check=True, - skip_one_per_day_check=True, ) ledger_collection.initial_load(client=None, sync=True) @@ -207,9 +199,7 @@ class TestBusiness: assert "Available Balance: 141" in res4 def test_addresses( - self, - gr_business: Business, - gr_db: PostgresConfig, + self, gr_business: Business, gr_db: PostgresConfig, gr_business_address ): from generalresearch.models.gr.business import BusinessAddress @@ -223,8 +213,8 @@ class TestBusiness: def test_teams( self, gr_business: Business, - team: Team, - team_manager: TeamManager, + gr_team: Team, + gr_team_manager: TeamManager, gr_db: PostgresConfig, ): assert gr_business.teams is None @@ -233,7 +223,7 @@ class TestBusiness: assert isinstance(gr_business.teams, list) assert len(gr_business.teams) == 0 - team_manager.add_business(team=team, business=gr_business) + gr_team_manager.add_business(team=gr_team, business=gr_business) assert len(gr_business.teams) == 0 gr_business.prefetch_teams(pg_config=gr_db) assert len(gr_business.teams) == 1 @@ -266,6 +256,7 @@ class TestBusiness: def test_bank_accounts( self, gr_business: Business, + gr_business_bank_account, gr_business_bank_account_manager: BusinessBankAccountManager, ): assert gr_business.products is None @@ -341,13 +332,8 @@ class TestBusiness: create_main_accounts() p = product_factory(business=gr_business) thl_ledger_manager.get_account_or_create_bp_wallet(product=p) - business_payout_event_manager.set_account_lookup_table( - thl_lm=thl_ledger_manager - ) - brokerage_product_payout_event_factory( - product=p, amount=USDCent(123), skip_wallet_balance_check=True - ) + brokerage_product_payout_event_factory(product=p, amount=USDCent(123)) gr_business.prebuild_payouts( bpem=business_payout_event_manager, @@ -359,18 +345,13 @@ class TestBusiness: brokerage_product_payout_event_factory( product=p, amount=USDCent(123), - skip_wallet_balance_check=True, - skip_one_per_day_check=True, - ) - business_payout_event_manager.set_account_lookup_table( - thl_lm=thl_ledger_manager ) gr_business.prebuild_payouts( bpem=business_payout_event_manager, ) assert isinstance(gr_business.payouts, list) - assert len(gr_business.payouts) == 1 - assert len(gr_business.payouts[0].bp_payouts) == 2 + assert len(gr_business.payouts) == 2 + assert len(gr_business.payouts[0].bp_payouts) == 1 assert sum([p.amount for p in gr_business.payouts]) == 246 def test_payouts_totals( @@ -390,29 +371,20 @@ class TestBusiness: p1: Product = product_factory(business=gr_business) thl_ledger_manager.get_account_or_create_bp_wallet(product=p1) - business_payout_event_manager.set_account_lookup_table( - thl_lm=thl_ledger_manager - ) brokerage_product_payout_event_factory( product=p1, amount=USDCent(1), - skip_wallet_balance_check=True, - skip_one_per_day_check=True, ) brokerage_product_payout_event_factory( product=p1, amount=USDCent(25), - skip_wallet_balance_check=True, - skip_one_per_day_check=True, ) brokerage_product_payout_event_factory( product=p1, amount=USDCent(50), - skip_wallet_balance_check=True, - skip_one_per_day_check=True, ) gr_business.prebuild_payouts( @@ -420,8 +392,10 @@ class TestBusiness: ) assert isinstance(gr_business.payouts, list) - assert len(gr_business.payouts) == 1 - assert len(gr_business.payouts[0].bp_payouts) == 3 + assert len(gr_business.payouts) == 3 + assert len(gr_business.payouts[0].bp_payouts) == 1 + assert len(gr_business.payouts[1].bp_payouts) == 1 + assert len(gr_business.payouts[2].bp_payouts) == 1 assert gr_business.payouts_total == USDCent(76) assert gr_business.payouts_total_str == "$0.76" @@ -467,7 +441,6 @@ class TestBusiness: class TestBusinessBalance: - @pytest.fixture def start(self) -> datetime: return datetime(year=2018, month=3, day=14, hour=0, tzinfo=UTC) @@ -678,16 +651,12 @@ class TestBusinessBalance: product=u1.product, amount=USDCent(5), created=start + timedelta(days=4), - skip_wallet_balance_check=True, - skip_one_per_day_check=True, ) brokerage_product_payout_event_factory( product=u2.product, amount=USDCent(50), created=start + timedelta(days=4), - skip_wallet_balance_check=True, - skip_one_per_day_check=True, ) ledger_collection.initial_load(client=None, sync=True) @@ -773,16 +742,12 @@ class TestBusinessBalance: product=u1.product, amount=USDCent(250), created=start + timedelta(days=3), - skip_wallet_balance_check=True, - skip_one_per_day_check=True, ) brokerage_product_payout_event_factory( product=u2.product, amount=USDCent(50), created=start + timedelta(days=4), - skip_wallet_balance_check=True, - skip_one_per_day_check=True, ) adj_to_fail_with_tx_factory(session=s1, created=start + timedelta(days=5)) @@ -889,8 +854,6 @@ class TestBusinessBalance: amount=USDCent(71), ext_ref_id=uuid4().hex, created=start + timedelta(days=1, minutes=1), - skip_wallet_balance_check=True, - skip_one_per_day_check=True, ) adj_to_fail_with_tx_factory( session=s1, @@ -1042,16 +1005,12 @@ class TestBusinessBalance: product=u1.product, amount=USDCent(250), created=start + timedelta(days=3), - skip_wallet_balance_check=True, - skip_one_per_day_check=True, ) brokerage_product_payout_event_factory( product=u2.product, amount=USDCent(50), created=start + timedelta(days=4), - skip_wallet_balance_check=True, - skip_one_per_day_check=True, ) session_with_tx_factory( @@ -1177,7 +1136,6 @@ class TestBusinessBalance: class TestBusinessMethods: - @pytest.fixture(scope="function") def start(self, utc_90days_ago: datetime) -> datetime: s = utc_90days_ago.replace(microsecond=0) @@ -1269,7 +1227,7 @@ class TestBusinessMethods: gr_db: PostgresConfig, thl_web_rr: PostgresConfig, product_factory: Callable[..., Product], - team: Team, + gr_team: Team, client_no_amm: DaskClient, mnt_filepath: GRLDatasets, ledger_manager: LedgerManager, @@ -1282,7 +1240,7 @@ class TestBusinessMethods: create_main_accounts: Callable[..., None], session_with_tx_factory: Callable[..., Session], ledger_collection, - team_manager: TeamManager, + gr_team_manager: TeamManager, pop_ledger_merge: PopLedgerMerge, gr_redis_config: RedisConfig, utc_60days_ago: datetime, @@ -1290,9 +1248,9 @@ class TestBusinessMethods: ): from generalresearch.models.gr.business import Business - p1 = product_factory(team=team, business=gr_business) + p1 = product_factory(team=gr_team, business=gr_business) u1 = user_factory(product=p1) - team_manager.add_business(team=team, business=gr_business) + gr_team_manager.add_business(team=gr_team, business=gr_business) # Business needs tx & incite to build balance delete_ledger_db() @@ -1345,7 +1303,7 @@ class TestBusinessMethods: assert isinstance(business2.teams, list) assert p1.uuid in [p.uuid for p in business2.products] assert len(business2.teams) == 1 - assert team.uuid in [t.uuid for t in business2.teams] + assert gr_team.uuid in [t.uuid for t in business2.teams] assert isinstance(business2.balance, BusinessBalances) assert business2.balance.payout == 48 diff --git a/tests/models/gr/test_team.py b/tests/models/gr/test_team.py index b5f1781..e853817 100644 --- a/tests/models/gr/test_team.py +++ b/tests/models/gr/test_team.py @@ -61,6 +61,7 @@ class TestTeam: self, gr_team: Team, gr_user: GRUser, + gr_membership, gr_user_factory: Callable[..., GRUser], gr_membership_manager: MembershipManager, ): @@ -105,7 +106,7 @@ class TestTeam: def test_businesses( self, gr_team: Team, - business: Business, + gr_business: Business, team_manager: TeamManager, gr_business_manager: BusinessManager, ): @@ -116,12 +117,12 @@ class TestTeam: assert isinstance(gr_team.businesses, list) assert len(gr_team.businesses) == 0 - team_manager.add_business(team=gr_team, business=business) + team_manager.add_business(team=gr_team, business=gr_business) assert len(gr_team.businesses) == 0 gr_team.prefetch_businesses(gr_business_manager=gr_business_manager) assert len(gr_team.businesses) == 1 assert isinstance(gr_team.businesses[0], Business) - assert gr_team.businesses[0].uuid == business.uuid + assert gr_team.businesses[0].uuid == gr_business.uuid def test_products( self, @@ -174,7 +175,6 @@ class TestTeamMethods: gr_user_manager=gr_user_manager, gr_business_manager=gr_business_manager, gr_membership_manager=gr_membership_manager, - thl_web_rr=thl_web_rr, redis_config=gr_redis_config, client=client_no_amm, ds=mnt_filepath, @@ -192,7 +192,7 @@ class TestTeamMethods: thl_web_rr: PostgresConfig, product_factory: Callable[..., Product], gr_team: Team, - membership_factory: Callable[..., Membership], + gr_membership_factory: Callable[..., Membership], gr_redis_config: RedisConfig, mnt_filepath: GRLDatasets, mnt_gr_api_dir: Path, @@ -206,14 +206,13 @@ class TestTeamMethods: from generalresearch.models.gr.team import Team p1 = product_factory(team=gr_team) - membership_factory(team=gr_team, gr_user=gr_user) + gr_membership_factory(gr_team=gr_team, gr_user=gr_user) gr_team.set_cache( product_manager=product_manager, gr_user_manager=gr_user_manager, gr_business_manager=gr_business_manager, gr_membership_manager=gr_membership_manager, - thl_web_rr=thl_web_rr, redis_config=gr_redis_config, client=client_no_amm, ds=mnt_filepath, diff --git a/tests/models/thl/test_contest/test_leaderboard_contest.py b/tests/models/thl/test_contest/test_leaderboard_contest.py index f787bdf..a639261 100644 --- a/tests/models/thl/test_contest/test_leaderboard_contest.py +++ b/tests/models/thl/test_contest/test_leaderboard_contest.py @@ -100,14 +100,14 @@ class TestLeaderboardContest(TestContest): def test_win( self, leaderboard_contest: LeaderboardContest, - thl_redis: Redis, + thl_redis_client: Redis, user_1: User, user_2: User, user_3: User, ): model = leaderboard_contest.leaderboard_model lbm = LeaderboardManager( - redis_client=thl_redis, + redis_client=thl_redis_client, board_code=model.board_code, country_iso=model.country_iso, freq=model.freq, diff --git a/tests/models/thl/test_product.py b/tests/models/thl/test_product.py index 25affcf..97abf0c 100644 --- a/tests/models/thl/test_product.py +++ b/tests/models/thl/test_product.py @@ -713,8 +713,6 @@ class TestProductFinancials: product=p1, amount=USDCent(50), created=start + timedelta(days=3), - skip_wallet_balance_check=True, - skip_one_per_day_check=True, ) assert ( len( @@ -754,7 +752,7 @@ class TestProductFinancials: ) assert p1.payouts is not None assert len(p1.payouts) == 1 - assert p1.payouts_total == 50 + assert p1.payouts_total == USDCent(50) assert p1.payouts_total_str == "$0.50" # -- Now pay ou another!. @@ -763,8 +761,6 @@ class TestProductFinancials: product=p1, amount=USDCent(5), created=start + timedelta(days=4), - skip_wallet_balance_check=True, - skip_one_per_day_check=True, ) assert ( len( @@ -862,8 +858,6 @@ class TestProductBalance: amount=USDCent(71), ext_ref_id=uuid4().hex, created=start + timedelta(days=1, minutes=1), - skip_wallet_balance_check=True, - skip_one_per_day_check=True, ) ledger_collection.initial_load(client=None, sync=True) pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) @@ -922,8 +916,6 @@ class TestProductBalance: amount=USDCent(71), ext_ref_id=uuid4().hex, created=datetime.now(tz=UTC), - skip_wallet_balance_check=True, - skip_one_per_day_check=True, ) ledger_collection.initial_load(client=None, sync=True) pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) @@ -1134,8 +1126,6 @@ class TestProductCache: amount=USDCent(71), ext_ref_id=uuid4().hex, created=start + timedelta(days=1, minutes=1), - skip_wallet_balance_check=True, - skip_one_per_day_check=True, ) # 3. Recon -- cgit v1.2.3