aboutsummaryrefslogtreecommitdiff
path: root/tests
diff options
context:
space:
mode:
authorstuppie2026-09-07 11:47:43 -0600
committerstuppie2026-09-07 11:47:43 -0600
commit092960233652cce1f4dc7841856034a6635e9cd9 (patch)
tree46e5fcd4d1e1b7ed0b987980c6c67ffa6e6b45c7 /tests
parent80fd8aab4c7271ddb619b0de18741d7ac77b490b (diff)
parent242579a44855873d5e054e375440e9d3492cd682 (diff)
downloadgeneralresearch-092960233652cce1f4dc7841856034a6635e9cd9.tar.gz
generalresearch-092960233652cce1f4dc7841856034a6635e9cd9.zip
Merge branch 'master' into dev-greg
Diffstat (limited to 'tests')
-rw-r--r--tests/conftest.py5
-rw-r--r--tests/grliq/managers/test_forensic_data.py65
-rw-r--r--tests/grliq/managers/test_forensic_results.py11
-rw-r--r--tests/grliq/models/test_forensic_data.py8
-rw-r--r--tests/incite/collections/test_df_collection_base.py38
-rw-r--r--tests/incite/collections/test_df_collection_item_base.py53
-rw-r--r--tests/incite/collections/test_df_collection_item_thl_web.py454
-rw-r--r--tests/incite/collections/test_df_collection_thl_marketplaces.py30
-rw-r--r--tests/incite/collections/test_df_collection_thl_web.py131
-rw-r--r--tests/incite/mergers/foundations/test_enriched_session.py76
-rw-r--r--tests/incite/mergers/foundations/test_enriched_task_adjust.py49
-rw-r--r--tests/incite/mergers/foundations/test_enriched_wall.py119
-rw-r--r--tests/incite/mergers/foundations/test_user_id_product.py47
-rw-r--r--tests/incite/mergers/test_merge_collection.py73
-rw-r--r--tests/incite/mergers/test_merge_collection_item.py37
-rw-r--r--tests/incite/mergers/test_pop_ledger.py148
-rw-r--r--tests/incite/mergers/test_ym_survey_merge.py84
-rw-r--r--tests/incite/schemas/test_admin_responses.py62
-rw-r--r--tests/incite/schemas/test_thl_web.py8
-rw-r--r--tests/incite/test_collection_base.py91
-rw-r--r--tests/incite/test_collection_base_item.py82
-rw-r--r--tests/incite/test_grl_flow.py11
-rw-r--r--tests/incite/test_interval_idx.py9
-rw-r--r--tests/managers/gr/test_authentication.py122
-rw-r--r--tests/managers/gr/test_business.py139
-rw-r--r--tests/managers/gr/test_team.py130
-rw-r--r--tests/managers/leaderboard.py177
-rw-r--r--tests/managers/network/__init__.py0
-rw-r--r--tests/managers/network/test_label.py202
-rw-r--r--tests/managers/network/test_tool_run.py25
-rw-r--r--tests/managers/test_events.py146
-rw-r--r--tests/managers/test_lucid.py9
-rw-r--r--tests/managers/test_userpid.py6
-rw-r--r--tests/managers/thl/test_buyer.py16
-rw-r--r--tests/managers/thl/test_cashout_method.py76
-rw-r--r--tests/managers/thl/test_category.py112
-rw-r--r--tests/managers/thl/test_contest/test_leaderboard.py89
-rw-r--r--tests/managers/thl/test_contest/test_milestone.py136
-rw-r--r--tests/managers/thl/test_contest/test_raffle.py232
-rw-r--r--tests/managers/thl/test_harmonized_uqa.py21
-rw-r--r--tests/managers/thl/test_ipinfo.py85
-rw-r--r--tests/managers/thl/test_ledger/test_lm_accounts.py162
-rw-r--r--tests/managers/thl/test_ledger/test_lm_tx.py145
-rw-r--r--tests/managers/thl/test_ledger/test_lm_tx_entries.py30
-rw-r--r--tests/managers/thl/test_ledger/test_lm_tx_locks.py283
-rw-r--r--tests/managers/thl/test_ledger/test_lm_tx_metadata.py43
-rw-r--r--tests/managers/thl/test_ledger/test_thl_lm_accounts.py329
-rw-r--r--tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py473
-rw-r--r--tests/managers/thl/test_ledger/test_thl_lm_tx.py1213
-rw-r--r--tests/managers/thl/test_ledger/test_thl_lm_tx__user_payouts.py404
-rw-r--r--tests/managers/thl/test_ledger/test_thl_pem.py139
-rw-r--r--tests/managers/thl/test_ledger/test_user_txs.py119
-rw-r--r--tests/managers/thl/test_ledger/test_wallet.py51
-rw-r--r--tests/managers/thl/test_maxmind.py500
-rw-r--r--tests/managers/thl/test_payout.py1073
-rw-r--r--tests/managers/thl/test_product.py119
-rw-r--r--tests/managers/thl/test_product_prod.py27
-rw-r--r--tests/managers/thl/test_profiling/test_question.py32
-rw-r--r--tests/managers/thl/test_profiling/test_schema.py18
-rw-r--r--tests/managers/thl/test_profiling/test_uqa.py1
-rw-r--r--tests/managers/thl/test_profiling/test_user_upk.py28
-rw-r--r--tests/managers/thl/test_session_manager.py94
-rw-r--r--tests/managers/thl/test_survey.py96
-rw-r--r--tests/managers/thl/test_survey_penalty.py27
-rw-r--r--tests/managers/thl/test_task_adjustment.py248
-rw-r--r--tests/managers/thl/test_task_status.py169
-rw-r--r--tests/managers/thl/test_user_manager/test_base.py76
-rw-r--r--tests/managers/thl/test_user_manager/test_mysql.py31
-rw-r--r--tests/managers/thl/test_user_manager/test_redis.py46
-rw-r--r--tests/managers/thl/test_user_manager/test_user_fetch.py19
-rw-r--r--tests/managers/thl/test_user_manager/test_user_metadata.py47
-rw-r--r--tests/managers/thl/test_user_streak.py73
-rw-r--r--tests/managers/thl/test_userhealth.py169
-rw-r--r--tests/managers/thl/test_wall_manager.py106
-rw-r--r--tests/models/admin/test_report_request.py20
-rw-r--r--tests/models/custom_types/test_aware_datetime.py8
-rw-r--r--tests/models/custom_types/test_dsn.py10
-rw-r--r--tests/models/custom_types/test_therest.py2
-rw-r--r--tests/models/dynata/test_eligbility.py8
-rw-r--r--tests/models/dynata/test_survey.py3
-rw-r--r--tests/models/gr/test_authentication.py222
-rw-r--r--tests/models/gr/test_base.py21
-rw-r--r--tests/models/gr/test_business.py1134
-rw-r--r--tests/models/gr/test_team.py356
-rw-r--r--tests/models/innovate/test_question.py10
-rw-r--r--tests/models/legacy/test_offerwall_parse_response.py4
-rw-r--r--tests/models/legacy/test_profiling_questions.py6
-rw-r--r--tests/models/legacy/test_user_question_answer_in.py69
-rw-r--r--tests/models/morning/test.py8
-rw-r--r--tests/models/network/__init__.py0
-rw-r--r--tests/models/network/test_mtr.py26
-rw-r--r--tests/models/network/test_nmap.py30
-rw-r--r--tests/models/network/test_nmap_parser.py22
-rw-r--r--tests/models/network/test_rdns.py34
-rw-r--r--tests/models/precision/__init__.py115
-rw-r--r--tests/models/precision/test_survey.py42
-rw-r--r--tests/models/prodege/test_survey_participation.py23
-rw-r--r--tests/models/spectrum/test_question.py21
-rw-r--r--tests/models/spectrum/test_survey.py86
-rw-r--r--tests/models/spectrum/test_survey_manager.py110
-rw-r--r--tests/models/test_currency.py126
-rw-r--r--tests/models/test_device.py8
-rw-r--r--tests/models/test_finance.py127
-rw-r--r--tests/models/thl/question/test_question_info.py139
-rw-r--r--tests/models/thl/question/test_user_info.py29
-rw-r--r--tests/models/thl/test_adjustments.py128
-rw-r--r--tests/models/thl/test_bucket.py8
-rw-r--r--tests/models/thl/test_buyer.py4
-rw-r--r--tests/models/thl/test_contest/test_contest.py10
-rw-r--r--tests/models/thl/test_contest/test_leaderboard_contest.py44
-rw-r--r--tests/models/thl/test_contest/test_raffle_contest.py44
-rw-r--r--tests/models/thl/test_ledger.py6
-rw-r--r--tests/models/thl/test_marketplace_condition.py42
-rw-r--r--tests/models/thl/test_payout.py120
-rw-r--r--tests/models/thl/test_payout_format.py2
-rw-r--r--tests/models/thl/test_product.py305
-rw-r--r--tests/models/thl/test_product_userwalletconfig.py10
-rw-r--r--tests/models/thl/test_soft_pair.py12
-rw-r--r--tests/models/thl/test_upkquestion.py86
-rw-r--r--tests/models/thl/test_user.py141
-rw-r--r--tests/models/thl/test_user_iphistory.py6
-rw-r--r--tests/models/thl/test_user_metadata.py4
-rw-r--r--tests/models/thl/test_user_streak.py4
-rw-r--r--tests/models/thl/test_wall.py42
-rw-r--r--tests/models/thl/test_wall_session.py24
-rw-r--r--tests/sql_helper.py4
-rw-r--r--tests/test_postgres.py39
-rw-r--r--tests/wall_status_codes/test_analyze.py2
-rw-r--r--tests/wxet/models/test_definitions.py37
-rw-r--r--tests/wxet/models/test_finish_type.py2
130 files changed, 7293 insertions, 6456 deletions
diff --git a/tests/conftest.py b/tests/conftest.py
index 6748592..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,9 @@ 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
+ "test_utils.precision.conftest",
+ "test_utils.spectrum.conftest",
]
diff --git a/tests/grliq/managers/test_forensic_data.py b/tests/grliq/managers/test_forensic_data.py
index e4854e8..2254829 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_factory(
+ 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/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..6d715fa 100644
--- a/tests/incite/collections/test_df_collection_base.py
+++ b/tests/incite/collections/test_df_collection_base.py
@@ -1,18 +1,18 @@
-from datetime import datetime, timezone
+from datetime import UTC, datetime
from typing import TYPE_CHECKING
import pandas as pd
import pytest
from pandera.pandas import DataFrameSchema
-from generalresearch.incite.collections import (
+from generalresearch.incite.collections.base import (
DFCollection,
DFCollectionType,
)
-from test_utils.incite.conftest import mnt_filepath
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]
@@ -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,24 +46,28 @@ 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=timezone.utc),
- finished=datetime(year=1900, month=1, day=1, tzinfo=timezone.utc),
- offset="100d",
+ 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),
)
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=timezone.utc),
- finished=datetime(year=1900, month=1, day=1, tzinfo=timezone.utc),
- offset="100d",
+ 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 +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
)
@@ -88,12 +94,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: PostgresConfig):
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..7a8793d 100644
--- a/tests/incite/collections/test_df_collection_item_base.py
+++ b/tests/incite/collections/test_df_collection_item_base.py
@@ -1,30 +1,30 @@
-from datetime import datetime, timezone
+from __future__ import annotations
+
+from datetime import UTC, datetime
from typing import TYPE_CHECKING
import pytest
-from generalresearch.incite.collections import (
+from generalresearch.incite.collections.base import (
+ MYSQL_ALLOWED_COLL_TYPES,
DFCollection,
DFCollectionItem,
DFCollectionType,
)
-from generalresearch.pg_helper import PostgresConfig
if TYPE_CHECKING:
from generalresearch.incite.base import GRLDatasets
-
-df_collection_types = [e for e in DFCollectionType if e is not DFCollectionType.TEST]
+ from generalresearch.pg_helper import PostgresConfig
-@pytest.mark.parametrize("df_coll_type", df_collection_types)
+@pytest.mark.parametrize("df_coll_type", MYSQL_ALLOWED_COLL_TYPES)
class TestDFCollectionItemBase:
-
- def test_init(self, mnt_filepath: "GRLDatasets", df_coll_type):
+ def test_init(self, mnt_filepath: GRLDatasets, df_coll_type: DFCollectionType):
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),
+ 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),
)
@@ -34,23 +34,23 @@ class TestDFCollectionItemBase:
assert isinstance(item, DFCollectionItem)
-@pytest.mark.parametrize("df_coll_type", df_collection_types)
+@pytest.mark.parametrize("df_coll_type", MYSQL_ALLOWED_COLL_TYPES)
class TestDFCollectionItemProperties:
-
@pytest.mark.skip
- def test_filename(self, df_coll_type):
+ def test_filename(self, df_coll_type: DFCollectionType):
pass
-@pytest.mark.parametrize("df_coll_type", df_collection_types)
+@pytest.mark.parametrize("df_coll_type", MYSQL_ALLOWED_COLL_TYPES)
class TestDFCollectionItemMethods:
-
- def test_has_mysql_false(self, mnt_filepath: "GRLDatasets", df_coll_type):
+ def test_has_mysql_false(
+ self, mnt_filepath: GRLDatasets, df_coll_type: DFCollectionType
+ ):
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),
+ 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),
)
@@ -58,13 +58,16 @@ 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,
- offset="100d",
- start=datetime(year=1800, month=1, day=1, tzinfo=timezone.utc),
- finished=datetime(year=1900, month=1, day=1, tzinfo=timezone.utc),
+ 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),
pg_config=thl_web_rr,
)
@@ -74,5 +77,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_item_thl_web.py b/tests/incite/collections/test_df_collection_item_thl_web.py
index 8b8bcbe..eeabb41 100644
--- a/tests/incite/collections/test_df_collection_item_thl_web.py
+++ b/tests/incite/collections/test_df_collection_item_thl_web.py
@@ -1,17 +1,19 @@
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
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
import pandas as pd
import pytest
-from distributed import Client, Scheduler, Worker
+from dask.distributed import Client as DaskClient
+from dask.distributed import Scheduler as DaskScheduler
+from dask.distributed import Worker as DaskWorker
# noinspection PyUnresolvedReferences
from distributed.utils_test import (
@@ -22,18 +24,21 @@ from pandera.pandas import DataFrameSchema
from pydantic import FilePath
from generalresearch.incite.base import CollectionItemBase
-from generalresearch.incite.collections import (
- DFCollectionItem,
+from generalresearch.incite.collections.base import (
DFCollectionType,
)
from generalresearch.incite.schemas import ARCHIVE_AFTER
-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.models.thl.product import Product
+ from generalresearch.models.thl.user import User
fake = Faker()
@@ -52,12 +57,11 @@ unsupported_mock_types = {
}
-def combo_object() -> Generator[str, None, None]:
- for x in iter_product(
+def combo_object() -> Generator[tuple[DFCollectionType, str]]:
+ yield from iter_product(
df_collections,
- ["15min", "45min", "1H"],
- ):
- yield from x
+ ["15min", "45min", "1h"],
+ )
class TestDFCollectionItemBase:
@@ -71,8 +75,12 @@ class TestDFCollectionItemBase:
argnames="df_collection_data_type, offset", argvalues=combo_object()
)
class TestDFCollectionItemProperties:
-
- def test_filename(self, df_collection_data_type, df_collection, offset: str):
+ def test_filename(
+ self,
+ df_collection_data_type: DFCollectionType,
+ df_collection: DFCollection,
+ offset: str,
+ ):
for i in df_collection.items:
assert isinstance(i.filename, str)
@@ -88,38 +96,59 @@ class TestDFCollectionItemProperties:
argnames="df_collection_data_type, offset", argvalues=combo_object()
)
class TestDFCollectionItemPropertiesBase:
-
- def test_name(self, df_collection_data_type, offset: str, df_collection):
+ def test_name(
+ self,
+ df_collection: DFCollection,
+ ):
for i in df_collection.items:
assert isinstance(i.name, str)
- def test_finish(self, df_collection_data_type, offset: str, df_collection):
+ def test_finish(
+ self,
+ df_collection: DFCollection,
+ ):
for i in df_collection.items:
assert isinstance(i.finish, datetime)
- def test_interval(self, df_collection_data_type, offset: str, df_collection):
+ def test_interval(
+ self,
+ df_collection: DFCollection,
+ ):
for i in df_collection.items:
assert isinstance(i.interval, pd.Interval)
def test_partial_filename(
- self, df_collection_data_type, offset: str, df_collection
+ self,
+ df_collection: DFCollection,
):
for i in df_collection.items:
assert isinstance(i.partial_filename, str)
- def test_empty_filename(self, df_collection_data_type, offset: str, df_collection):
+ def test_empty_filename(
+ self,
+ df_collection: DFCollection,
+ ):
for i in df_collection.items:
assert isinstance(i.empty_filename, str)
- def test_path(self, df_collection_data_type, offset: str, df_collection):
+ def test_path(
+ self,
+ df_collection: DFCollection,
+ ):
for i in df_collection.items:
assert isinstance(i.path, FilePath)
- def test_partial_path(self, df_collection_data_type, offset: str, df_collection):
+ def test_partial_path(
+ self,
+ df_collection: DFCollection,
+ ):
for i in df_collection.items:
assert isinstance(i.partial_path, FilePath)
- def test_empty_path(self, df_collection_data_type, offset: str, df_collection):
+ def test_empty_path(
+ self,
+ df_collection: DFCollection,
+ ):
for i in df_collection.items:
assert isinstance(i.empty_path, FilePath)
@@ -135,26 +164,25 @@ class TestDFCollectionItemPropertiesBase:
),
)
class TestDFCollectionItemMethod:
-
- def test_has_mysql(
+ def test_has_postgres(
self,
- df_collection,
- thl_web_rr: PostgresConfig,
+ df_collection_data_type: DFCollectionType,
offset: str,
duration: timedelta,
- df_collection_data_type,
- delete_df_collection,
+ delete_df_collection: Callable[..., None],
+ df_collection: DFCollection,
+ thl_web_rr: PostgresConfig,
):
delete_df_collection(coll=df_collection)
df_collection.pg_config = None
for i in df_collection.items:
- assert not i.has_mysql()
+ assert not i.has_postgres()
# Confirm that the regular connection should work as expected
df_collection.pg_config = thl_web_rr
for i in df_collection.items:
- assert i.has_mysql()
+ assert i.has_postgres()
# Make a fake connection and confirm it does NOT work
df_collection.pg_config = PostgresConfig(
@@ -163,17 +191,14 @@ class TestDFCollectionItemMethod:
statement_timeout=1,
)
for i in df_collection.items:
- assert not i.has_mysql()
+ assert not i.has_postgres()
@pytest.mark.skip
def test_update_partial_archive(
self,
- df_collection,
+ df_collection_data_type: DFCollectionType,
offset: str,
duration: timedelta,
- thl_web_rw: PostgresConfig,
- df_collection_data_type,
- delete_df_collection,
):
# for i in collection.items:
# assert i.update_partial_archive()
@@ -183,29 +208,16 @@ class TestDFCollectionItemMethod:
@pytest.mark.skip
def test_create_partial_archive(
self,
- df_collection,
+ df_collection_data_type: DFCollectionType,
offset: str,
- duration: str,
- create_main_accounts,
- thl_web_rw: PostgresConfig,
- thl_lm,
- df_collection_data_type,
- user_factory: Callable[..., User],
- product: Product,
- client_no_amm,
- incite_item_factory,
- delete_df_collection,
- mnt_filepath: GRLDatasets,
+ duration: timedelta,
):
- assert 1 + 1 == 2
+ pass
def test_dict(
self,
- df_collection_data_type,
- offset: str,
- duration: timedelta,
- df_collection,
- delete_df_collection,
+ df_collection: DFCollection,
+ delete_df_collection: Callable[..., None],
):
delete_df_collection(coll=df_collection)
@@ -225,18 +237,17 @@ class TestDFCollectionItemMethod:
def test_from_mysql(
self,
- df_collection_data_type,
- df_collection,
+ df_collection_data_type: DFCollectionType,
+ df_collection: DFCollection,
offset: str,
duration: timedelta,
- create_main_accounts,
+ create_main_accounts: Callable[..., None],
thl_web_rw: PostgresConfig,
user_factory: Callable[..., User],
product: Product,
- incite_item_factory,
- delete_df_collection,
+ incite_item_factory: Callable[..., None],
+ delete_df_collection: Callable[..., None],
):
- from generalresearch.models.thl.user import User
if df_collection.data_type in unsupported_mock_types:
return
@@ -249,38 +260,32 @@ class TestDFCollectionItemMethod:
for item in df_collection.items:
# Unlike .from_mysql_ledger(), .from_mysql_standard() will return
# back and empty df with the correct columns in place
- delete_df_collection(coll=df_collection)
- df = item.from_mysql()
if df_collection.data_type == DFCollectionType.LEDGER:
- assert df is None
- else:
- assert df.empty
- assert set(df.columns) == set(df_collection._schema.columns.keys())
+ continue
+ delete_df_collection(coll=df_collection)
+ df = item.from_postgres_standard()
+ assert isinstance(df, pd.DataFrame)
+ assert df.empty
+ assert set(df.columns) == set(df_collection.type_schema.columns.keys())
incite_item_factory(user=u1, item=item)
- df = item.from_mysql()
+ df = item.from_postgres_standard()
+ assert isinstance(df, pd.DataFrame)
assert not df.empty
- assert set(df.columns) == set(df_collection._schema.columns.keys())
- if df_collection.data_type == DFCollectionType.LEDGER:
- # The number of rows in this dataframe will change depending
- # on the mocking of data. It's because if the account has
- # user wallet on, then there will be more transactions for
- # example.
- assert df.shape[0] > 0
+ assert set(df.columns) == set(df_collection.type_schema.columns.keys())
- def test_from_mysql_standard(
+ def test_from_postgres_standard(
self,
- df_collection_data_type,
- df_collection,
+ df_collection_data_type: DFCollectionType,
+ df_collection: DFCollection,
offset: str,
duration: timedelta,
user_factory: Callable[..., User],
product: Product,
- incite_item_factory,
- delete_df_collection,
+ incite_item_factory: Callable[..., None],
+ delete_df_collection: Callable[..., None],
):
- from generalresearch.models.thl.user import User
if df_collection.data_type in unsupported_mock_types:
return
@@ -292,48 +297,31 @@ class TestDFCollectionItemMethod:
item: DFCollectionItem
if df_collection.data_type == DFCollectionType.LEDGER:
- # We're using parametrize, so this If statement is just to
- # confirm other Item Types will always raise an assertion
- with pytest.raises(expected_exception=AssertionError) as cm:
- res = item.from_mysql_standard()
- assert (
- "Can't call from_mysql_standard for Ledger DFCollectionItem"
- in str(cm.value)
- )
-
continue
# Unlike .from_mysql_ledger(), .from_mysql_standard() will return
# back and empty df with the correct columns in place
- df = item.from_mysql_standard()
+ df = item.from_postgres_standard()
+ assert isinstance(df, pd.DataFrame)
assert df.empty
- assert set(df.columns) == set(df_collection._schema.columns.keys())
+ assert set(df.columns) == set(df_collection.type_schema.columns.keys())
incite_item_factory(user=u1, item=item)
- df = item.from_mysql_standard()
+ df = item.from_postgres_standard()
+ assert isinstance(df, pd.DataFrame)
assert not df.empty
- assert set(df.columns) == set(df_collection._schema.columns.keys())
+ assert set(df.columns) == set(df_collection.type_schema.columns.keys())
assert df.shape[0] > 0
- def test_from_mysql_ledger(
+ def test_from_postgres_ledger(
self,
- df_collection,
- user: User,
- create_main_accounts,
- offset: str,
- duration: timedelta,
- thl_web_rw: PostgresConfig,
- thl_lm,
- df_collection_data_type,
+ df_collection: DFCollection,
user_factory: Callable[..., User],
product: Product,
- client_no_amm,
- incite_item_factory,
- delete_df_collection,
- mnt_filepath,
+ incite_item_factory: Callable[..., None],
+ delete_df_collection: Callable[..., None],
):
- from generalresearch.models.thl.user import User
if df_collection.data_type != DFCollectionType.LEDGER:
return
@@ -348,14 +336,14 @@ class TestDFCollectionItemMethod:
# Okay, now continue with the actual Ledger Item tests... we need
# to ensure that this item.start - item.finish range hasn't had
# any prior transactions created within that range.
- assert item.from_mysql_ledger() is None
+ assert item.from_postgres_ledger() is None
# Create main accounts doesn't matter because it doesn't
# add any transactions to the db
- assert item.from_mysql_ledger() is None
+ assert item.from_postgres_ledger() is None
incite_item_factory(user=u1, item=item)
- df = item.from_mysql_ledger()
+ df = item.from_postgres_ledger()
assert isinstance(df, pd.DataFrame)
# Not only is this a np.int64 to int comparison, but I also know it
@@ -373,19 +361,12 @@ class TestDFCollectionItemMethod:
def test_to_archive(
self,
- df_collection,
- user: User,
- offset: str,
- duration: timedelta,
- df_collection_data_type,
+ df_collection: DFCollection,
user_factory: Callable[..., User],
product: Product,
- client_no_amm,
- incite_item_factory,
- delete_df_collection,
- mnt_filepath,
+ incite_item_factory: Callable[..., None],
+ delete_df_collection: Callable[..., None],
):
- from generalresearch.models.thl.user import User
if df_collection.data_type in unsupported_mock_types:
return
@@ -400,7 +381,7 @@ class TestDFCollectionItemMethod:
# Load up the data that we'll be using for various to_archive
# methods.
- df = item.from_mysql()
+ df = item.from_postgres_standard()
ddf = dd.from_pandas(df, npartitions=1)
# (1) Write the basic archive, the issue is that because it's
@@ -411,17 +392,12 @@ class TestDFCollectionItemMethod:
def test__to_archive(
self,
- df_collection_data_type,
- df_collection,
+ df_collection: DFCollection,
user_factory: Callable[..., User],
product: Product,
- offset: str,
- duration: timedelta,
- client_no_amm,
- user: User,
- incite_item_factory,
- delete_df_collection,
- mnt_filepath,
+ incite_item_factory: Callable[..., None],
+ delete_df_collection: Callable[..., None],
+ mnt_filepath: GRLDatasets,
):
"""We already have a test for the "non-private" version of this,
which primarily just uses the respective Client to determine if
@@ -443,7 +419,7 @@ class TestDFCollectionItemMethod:
# Load up the data that we'll be using for various to_archive
# methods. Will always be empty pd.DataFrames for now...
- df = item.from_mysql()
+ df = item.from_db()
ddf = dd.from_pandas(df, npartitions=1)
# (1) Confirm a missing ddf (shouldn't bc of type hint) should
@@ -484,19 +460,28 @@ class TestDFCollectionItemMethod:
@pytest.mark.skip
def test_to_archive_numbered_partial(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ duration: timedelta,
):
pass
@pytest.mark.skip
def test_initial_load(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ duration: timedelta,
):
pass
@pytest.mark.skip
def test_clear_corrupt_archive(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ duration: timedelta,
):
pass
@@ -506,37 +491,36 @@ class TestDFCollectionItemMethod:
argvalues=list(iter_product(df_collections, ["12h", "10D"], [timedelta(days=15)])),
)
class TestDFCollectionItemMethodBase:
-
- @pytest.mark.skip
- def test_path_exists(
- self, df_collection_data_type, offset: str, duration: timedelta
- ):
- pass
-
- @pytest.mark.skip
- def test_next_numbered_path(
- self, df_collection_data_type, offset: str, duration: timedelta
- ):
- pass
-
@pytest.mark.skip
def test_search_highest_numbered_path(
- self, df_collection_data_type, offset: str, duration: timedelta
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ duration: timedelta,
):
pass
@pytest.mark.skip
def test_tmp_filename(
- self, df_collection_data_type, offset: str, duration: timedelta
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ duration: timedelta,
):
pass
@pytest.mark.skip
- def test_tmp_path(self, df_collection_data_type, offset: str, duration: timedelta):
+ def test_tmp_path(
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ duration: timedelta,
+ ):
pass
def test_is_empty(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection: DFCollection,
):
"""
test_has_empty was merged into this because item.has_empty is
@@ -553,7 +537,8 @@ class TestDFCollectionItemMethodBase:
assert item.has_empty()
def test_has_partial_archive(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection: DFCollection,
):
for item in df_collection.items:
assert not item.has_partial_archive()
@@ -561,7 +546,8 @@ class TestDFCollectionItemMethodBase:
assert item.has_partial_archive()
def test_has_archive(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection: DFCollection,
):
for item in df_collection.items:
# (1) Originally, nothing exists... so let's just make a file and
@@ -598,7 +584,8 @@ class TestDFCollectionItemMethodBase:
assert item.has_archive(include_empty=True)
def test_delete_archive(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection: DFCollection,
):
for item in df_collection.items:
item: DFCollectionItem
@@ -621,9 +608,11 @@ class TestDFCollectionItemMethodBase:
assert not item.partial_path.exists()
def test_should_archive(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection: DFCollection,
):
- schema: DataFrameSchema = df_collection._schema
+ schema: DataFrameSchema = df_collection.type_schema
+ assert schema.metadata
aa = schema.metadata[ARCHIVE_AFTER]
# It shouldn't be None, it can be timedelta(seconds=0)
@@ -632,19 +621,23 @@ 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()
@pytest.mark.skip
def test_set_empty(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ duration: timedelta,
):
pass
def test_valid_archive(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection: DFCollection,
):
# Originally, nothing has been saved or anything.. so confirm it
# always comes back as None
@@ -668,18 +661,28 @@ class TestDFCollectionItemMethodBase:
@pytest.mark.skip
def test_validate_df(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ duration: timedelta,
):
pass
@pytest.mark.skip
def test_from_archive(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ duration: timedelta,
):
pass
def test__to_dict(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ duration: timedelta,
+ df_collection: DFCollection,
):
for item in df_collection.items:
@@ -698,30 +701,39 @@ class TestDFCollectionItemMethodBase:
@pytest.mark.skip
def test_delete_partial(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ duration: timedelta,
):
pass
@pytest.mark.skip
def test_cleanup_partials(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ duration: timedelta,
):
pass
@pytest.mark.skip
def test_delete_dangling_partials(
- self, df_collection_data_type, df_collection, offset: str, duration: timedelta
+ self,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+ duration: timedelta,
):
pass
@gen_cluster(client=True, nthreads=[("127.0.0.1", 1)])
-async def test_client(client, s, worker):
+async def test_client(client: DaskClient, s: DaskScheduler, worker: DaskWorker):
"""c,s,a are all required - the secondary Worker (b) is not required"""
- assert isinstance(client, Client)
- assert isinstance(s, Scheduler)
- assert isinstance(worker, Worker)
+ assert isinstance(client, DaskClient)
+ assert isinstance(s, DaskScheduler)
+ assert isinstance(worker, DaskWorker)
@pytest.mark.parametrize(
@@ -730,12 +742,18 @@ async def test_client(client, s, worker):
)
@gen_cluster(client=True, nthreads=[("127.0.0.1", 1)])
@pytest.mark.anyio
-async def test_client_parametrize(c, s, w, df_collection_data_type, offset: str):
+async def test_client_parametrize(
+ c: DaskClient,
+ s: DaskScheduler,
+ w: DaskWorker,
+ df_collection_data_type: DFCollectionType,
+ offset: str,
+):
"""c,s,a are all required - the secondary Worker (b) is not required"""
- assert isinstance(c, Client), f"c is not Client, it's {type(c)}"
- assert isinstance(s, Scheduler), f"s is not Scheduler, it's {type(s)}"
- assert isinstance(w, Worker), f"w is not Worker, it's {type(w)}"
+ assert isinstance(c, DaskClient), f"c is not Client, it's {type(c)}"
+ assert isinstance(s, DaskScheduler), f"s is not Scheduler, it's {type(s)}"
+ assert isinstance(w, DaskWorker), f"w is not Worker, it's {type(w)}"
assert df_collection_data_type is not None
assert isinstance(offset, str)
@@ -751,22 +769,15 @@ async def test_client_parametrize(c, s, w, df_collection_data_type, offset: str)
argvalues=list(iter_product(df_collections, ["12h", "10D"], [timedelta(days=15)])),
)
class TestDFCollectionItemFunctionalTest:
-
def test_to_archive_and_ddf(
self,
- df_collection_data_type,
- offset: str,
- duration: timedelta,
- client_no_amm,
- df_collection,
- user: User,
+ client_no_amm: DaskClient,
+ df_collection: DFCollection,
user_factory: Callable[..., User],
product: Product,
- incite_item_factory,
- delete_df_collection,
- mnt_filepath: GRLDatasets,
+ incite_item_factory: Callable[..., None],
+ delete_df_collection: Callable[..., None],
):
- from generalresearch.models.thl.user import User
if df_collection.data_type in unsupported_mock_types:
return
@@ -804,17 +815,11 @@ class TestDFCollectionItemFunctionalTest:
def test_filesize_estimate(
self,
- df_collection,
- user: User,
- offset: str,
- duration: timedelta,
- client_no_amm,
+ df_collection: DFCollection,
user_factory: Callable[..., User],
product: Product,
- df_collection_data_type,
- incite_item_factory,
- delete_df_collection,
- mnt_filepath: GRLDatasets,
+ incite_item_factory: Callable[..., None],
+ delete_df_collection: Callable[..., None],
):
"""A functional test to write some Parquet files for the
DFCollection and then confirm that the files get written
@@ -828,8 +833,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
delete_df_collection(coll=df_collection)
@@ -853,18 +856,13 @@ class TestDFCollectionItemFunctionalTest:
def test_to_archive_client(
self,
- client_no_amm,
- df_collection,
+ client_no_amm: DaskClient,
+ df_collection: DFCollection,
user_factory: Callable[..., User],
product: Product,
- offset: str,
- duration: timedelta,
- df_collection_data_type,
- incite_item_factory,
- delete_df_collection,
- mnt_filepath: GRLDatasets,
+ incite_item_factory: Callable[..., None],
+ delete_df_collection: Callable[..., None],
):
- from generalresearch.models.thl.user import User
delete_df_collection(coll=df_collection)
df_collection._client = client_no_amm
@@ -880,7 +878,7 @@ class TestDFCollectionItemFunctionalTest:
# Load up the data that we'll be using for various to_archive
# methods. Will always be empty pd.DataFrames for now...
- df = item.from_mysql()
+ df = item.from_db()
ddf = dd.from_pandas(df, npartitions=1)
assert isinstance(ddf, dd.DataFrame)
@@ -893,7 +891,8 @@ class TestDFCollectionItemFunctionalTest:
@pytest.mark.skip
def test_get_items(
- self, df_collection, product: Product, offset: str, duration: timedelta
+ self,
+ df_collection: DFCollection,
):
with pytest.warns(expected_warning=ResourceWarning) as cm:
df_collection.get_items_last365()
@@ -906,27 +905,22 @@ class TestDFCollectionItemFunctionalTest:
def test_saving_protections(
self,
- client_no_amm,
- df_collection_data_type,
- df_collection,
- incite_item_factory,
- delete_df_collection,
+ df_collection: DFCollection,
+ incite_item_factory: Callable[..., None],
+ delete_df_collection: Callable[..., None],
user_factory: Callable[..., User],
product: Product,
- offset: str,
- duration: timedelta,
- mnt_filepath: GRLDatasets,
):
"""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
u1: User = user_factory(product=product)
- schema: DataFrameSchema = df_collection._schema
+ schema: DataFrameSchema = df_collection.type_schema
+ assert schema.metadata
aa = schema.metadata[ARCHIVE_AFTER]
assert isinstance(aa, timedelta)
@@ -948,21 +942,14 @@ class TestDFCollectionItemFunctionalTest:
def test_empty_item(
self,
- client_no_amm,
- df_collection_data_type,
- df_collection,
- incite_item_factory,
- delete_df_collection,
- user: User,
- offset: str,
- duration: timedelta,
- mnt_filepath: GRLDatasets,
+ df_collection: DFCollection,
+ delete_df_collection: Callable[..., None],
):
delete_df_collection(coll=df_collection)
for item in df_collection.items:
assert not item.has_empty()
- df: pd.DataFrame = item.from_mysql()
+ df: pd.DataFrame = item.from_db()
# We do this check b/c the Ledger returns back None and
# I don't want it to fail when we go to make a ddf
@@ -976,18 +963,13 @@ class TestDFCollectionItemFunctionalTest:
def test_file_touching(
self,
- client_no_amm,
- df_collection_data_type,
- df_collection,
- incite_item_factory,
- delete_df_collection,
+ client_no_amm: DaskClient,
+ df_collection: DFCollection,
+ incite_item_factory: Callable[..., None],
+ delete_df_collection: Callable[..., None],
user_factory: Callable[..., User],
product: Product,
- offset: str,
- 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 981f62e..0f79b81 100644
--- a/tests/incite/collections/test_df_collection_thl_marketplaces.py
+++ b/tests/incite/collections/test_df_collection_thl_marketplaces.py
@@ -1,25 +1,26 @@
-from datetime import datetime, timezone
+from collections.abc import Generator
+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.collections import DFCollection, DFCollectionType
+from generalresearch.incite.collections.base import DFCollection, DFCollectionType
from generalresearch.incite.collections.thl_marketplaces import (
InnovateSurveyHistoryCollection,
MorningSurveyTimeseriesCollection,
SagoSurveyHistoryCollection,
SpectrumSurveyTimeseriesCollection,
)
-from test_utils.incite.conftest import mnt_filepath
if TYPE_CHECKING:
from generalresearch.incite.base import GRLDatasets
+ from generalresearch.pg_helper import PostgresConfig
-def combo_object():
- for x in product(
+def combo_object() -> Generator[tuple[type, str]]:
+ yield from product(
[
InnovateSurveyHistoryCollection,
MorningSurveyTimeseriesCollection,
@@ -27,14 +28,19 @@ def combo_object():
SpectrumSurveyTimeseriesCollection,
],
["5min", "6H", "30D"],
- ):
- yield from x
+ )
@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
@@ -43,7 +49,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=ValueError):
instance = df_coll()
# (2) Confirm it only needs the archive_path
@@ -57,8 +63,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)
@@ -66,7 +72,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 b09d44c..7253dd0 100644
--- a/tests/incite/collections/test_df_collection_thl_web.py
+++ b/tests/incite/collections/test_df_collection_thl_web.py
@@ -3,25 +3,20 @@ 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.base import GRLDatasets
- from generalresearch.incite.collections import (
- DFCollectionItem,
- DFCollectionType,
- )
+from generalresearch.incite.collections.base import (
+ DFCollection,
+ DFCollectionType,
+)
-def combo_object() -> Generator[tuple, None, None]:
- for x in product(
+def combo_object() -> Generator[tuple[DFCollectionType, str]]:
+ yield from product(
[
DFCollectionType.USER,
DFCollectionType.WALL,
@@ -30,9 +25,8 @@ def combo_object() -> Generator[tuple, None, None]:
DFCollectionType.AUDIT_LOG,
DFCollectionType.LEDGER,
],
- ["30min", "1H"],
- ):
- yield from x
+ ["30min", "1h"],
+ )
@pytest.mark.parametrize(
@@ -41,7 +35,10 @@ def combo_object() -> Generator[tuple, None, None]:
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)
@@ -52,12 +49,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)
@@ -67,16 +64,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)
@@ -86,17 +83,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
@@ -105,63 +106,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
@@ -170,18 +216,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_enriched_session.py b/tests/incite/mergers/foundations/test_enriched_session.py
index 47f243e..71b2442 100644
--- a/tests/incite/mergers/foundations/test_enriched_session.py
+++ b/tests/incite/mergers/foundations/test_enriched_session.py
@@ -1,20 +1,35 @@
-from datetime import datetime, timedelta, timezone
+from __future__ import annotations
+
+from collections.abc import Callable
+from datetime import UTC, datetime, timedelta
from decimal import Decimal
from itertools import product
-from typing import Optional
+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.schemas.admin_responses import (
AdminPOPSessionSchema,
)
-from generalresearch.pg_helper import PostgresConfig
-from test_utils.incite.collections.conftest import (
- session_collection,
- wall_collection,
-)
+
+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(
@@ -30,17 +45,16 @@ class TestEnrichedSession:
def test_base(
self,
- client_no_amm,
- product,
- user_factory,
- wall_collection,
- session_collection,
- enriched_session_merge,
+ client_no_amm: DaskClient,
+ product: Product,
+ user_factory: Callable[..., User],
+ wall_collection: WallDFCollection,
+ session_collection: SessionDFCollection,
+ enriched_session_merge: EnrichedSessionMerge,
thl_web_rr: PostgresConfig,
- delete_df_collection,
- incite_item_factory,
+ delete_df_collection: Callable[..., None],
+ incite_item_factory: Callable[..., None],
):
- from generalresearch.models.thl.user import User
delete_df_collection(coll=session_collection)
@@ -77,31 +91,31 @@ 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(
self,
- event_report_request,
- enriched_session_merge,
- client_no_amm,
- wall_collection,
- session_collection,
+ event_report_request: ReportRequest,
+ enriched_session_merge: EnrichedSessionMerge,
+ client_no_amm: DaskClient,
+ wall_collection: WallDFCollection,
+ session_collection: SessionDFCollection,
thl_web_rr: PostgresConfig,
- session_report_request,
- user_factory,
- start,
- session_factory,
- product_factory,
- delete_df_collection,
+ session_report_request: ReportRequest,
+ user_factory: Callable[..., User],
+ start: datetime,
+ session_factory: Callable[..., Session],
+ product_factory: Callable[..., Product],
+ delete_df_collection: Callable[..., None],
):
delete_df_collection(coll=wall_collection)
delete_df_collection(coll=session_collection)
@@ -112,7 +126,7 @@ class TestEnrichedSessionAdmin:
for p in [p1, p2]:
u = user_factory(product=p)
for i in range(50):
- s = session_factory(
+ _ = session_factory(
user=u,
wall_count=1,
wall_req_cpi=Decimal("1.00"),
diff --git a/tests/incite/mergers/foundations/test_enriched_task_adjust.py b/tests/incite/mergers/foundations/test_enriched_task_adjust.py
index 96c214f..877d22f 100644
--- a/tests/incite/mergers/foundations/test_enriched_task_adjust.py
+++ b/tests/incite/mergers/foundations/test_enriched_task_adjust.py
@@ -1,16 +1,30 @@
+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 test_utils.incite.collections.conftest import (
- wall_collection,
- task_adj_collection,
- session_collection,
-)
-from test_utils.incite.mergers.conftest import enriched_wall_merge
+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(
@@ -27,19 +41,18 @@ class TestEnrichedTaskAdjust:
@pytest.mark.skip
def test_base(
self,
- client_no_amm,
- user_factory,
- product,
- task_adj_collection,
- wall_collection,
- session_collection,
- enriched_wall_merge,
- enriched_task_adjust_merge,
- incite_item_factory,
- delete_df_collection,
- thl_web_rr,
+ client_no_amm: DaskClient,
+ user_factory: Callable[..., User],
+ product: Product,
+ task_adj_collection: TaskAdjustmentDFCollection,
+ wall_collection: WallDFCollection,
+ session_collection: SessionDFCollection,
+ enriched_wall_merge: EnrichedWallMerge,
+ enriched_task_adjust_merge: EnrichedTaskAdjustMerge,
+ incite_item_factory: Callable[..., None],
+ delete_df_collection: Callable[..., None],
+ thl_web_rr: PostgresConfig,
):
- from generalresearch.models.thl.user import User
# -- Build & Setup
delete_df_collection(coll=session_collection)
diff --git a/tests/incite/mergers/foundations/test_enriched_wall.py b/tests/incite/mergers/foundations/test_enriched_wall.py
index 8f4995b..2b9afb8 100644
--- a/tests/incite/mergers/foundations/test_enriched_wall.py
+++ b/tests/incite/mergers/foundations/test_enriched_wall.py
@@ -1,34 +1,33 @@
-from datetime import timedelta, timezone, datetime
+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 Optional
+from typing import TYPE_CHECKING
import dask.dataframe as dd
import pandas as pd
import pytest
-
-# noinspection PyUnresolvedReferences
-from distributed.utils_test import (
- gen_cluster,
- client_no_amm,
- loop,
- loop_in_thread,
- cleanup,
- cluster_fixture,
- client,
-)
+from dask.distributed import Client as DaskClient
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,
-)
+
+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(
@@ -39,17 +38,16 @@ class TestEnrichedWall:
def test_base(
self,
- client_no_amm,
- product,
- user_factory,
- wall_collection,
- thl_web_rr,
- session_collection,
- enriched_wall_merge,
- delete_df_collection,
- incite_item_factory,
+ client_no_amm: DaskClient,
+ product: Product,
+ user_factory: Callable[..., User],
+ wall_collection: WallDFCollection,
+ thl_web_rr: PostgresConfig,
+ session_collection: SessionDFCollection,
+ enriched_wall_merge: EnrichedWallMerge,
+ delete_df_collection: Callable[..., None],
+ incite_item_factory: Callable[..., None],
):
- from generalresearch.models.thl.user import User
# -- Build & Setup
delete_df_collection(coll=session_collection)
@@ -82,15 +80,15 @@ class TestEnrichedWall:
def test_base_item(
self,
- client_no_amm,
- product,
- user_factory,
- wall_collection,
- session_collection,
- enriched_wall_merge,
- delete_df_collection,
- thl_web_rr,
- incite_item_factory,
+ client_no_amm: DaskClient,
+ product: Product,
+ user_factory: Callable[..., User],
+ wall_collection: WallDFCollection,
+ session_collection: SessionDFCollection,
+ enriched_wall_merge: EnrichedWallMerge,
+ delete_df_collection: Callable[..., None],
+ thl_web_rr: PostgresConfig,
+ incite_item_factory: Callable[..., None],
):
# -- Build & Setup
delete_df_collection(coll=session_collection)
@@ -118,7 +116,7 @@ class TestEnrichedWall:
try:
modified_time1 = path.stat().st_mtime
- except (Exception,):
+ except OSError:
modified_time1 = 0
item.build(
@@ -158,18 +156,23 @@ 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):
+ def test_empty(
+ self,
+ enriched_wall_merge: EnrichedWallMerge,
+ client_no_amm: DaskClient,
+ start: datetime,
+ ):
from generalresearch.models.admin.request import ReportRequest
rr = ReportRequest.model_validate({"interval": "5min", "start": start})
@@ -186,18 +189,18 @@ class TestEnrichedWallToAdmin:
def test_to_admin_response(
self,
- event_report_request,
- enriched_wall_merge,
- client_no_amm,
- wall_collection,
- session_collection,
- thl_web_rr,
- user,
- session_factory,
- delete_df_collection,
- product_factory,
- user_factory,
- start,
+ event_report_request: ReportRequest,
+ enriched_wall_merge: EnrichedWallMerge,
+ client_no_amm: DaskClient,
+ wall_collection: WallDFCollection,
+ session_collection: SessionDFCollection,
+ thl_web_rr: PostgresConfig,
+ user: User,
+ session_factory: Callable[..., Session],
+ delete_df_collection: Callable[..., None],
+ product_factory: Callable[..., Product],
+ user_factory: Callable[..., User],
+ start: datetime,
):
delete_df_collection(coll=wall_collection)
delete_df_collection(coll=session_collection)
@@ -208,7 +211,7 @@ class TestEnrichedWallToAdmin:
for p in [p1, p2]:
u = user_factory(product=p)
for i in range(50):
- s = session_factory(
+ _ = session_factory(
user=u,
wall_count=2,
wall_req_cpi=Decimal("1.00"),
diff --git a/tests/incite/mergers/foundations/test_user_id_product.py b/tests/incite/mergers/foundations/test_user_id_product.py
index f96bfb4..8c4b2f7 100644
--- a/tests/incite/mergers/foundations/test_user_id_product.py
+++ b/tests/incite/mergers/foundations/test_user_id_product.py
@@ -1,24 +1,22 @@
-from datetime import timedelta, datetime, timezone
+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
-
-# noinspection PyUnresolvedReferences
-from distributed.utils_test import (
- gen_cluster,
- client_no_amm,
- loop,
- loop_in_thread,
- cleanup,
- cluster_fixture,
- client,
-)
+from dask.distributed import Client as DaskClient
from generalresearch.incite.mergers.foundations.user_id_product import (
UserIdProductMergeItem,
)
-from test_utils.incite.mergers.conftest import user_id_product_merge
+
+if TYPE_CHECKING:
+ # noinspection PyUnresolvedReferences
+ from generalresearch.incite.mergers.foundations.user_id_product import (
+ UserIdProductMerge,
+ )
@pytest.mark.parametrize(
@@ -27,25 +25,28 @@ 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)],
)
),
)
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:
@@ -55,7 +56,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)
@@ -64,7 +65,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_merge_collection.py b/tests/incite/mergers/test_merge_collection.py
index ec507bc..7ed3996 100644
--- a/tests/incite/mergers/test_merge_collection.py
+++ b/tests/incite/mergers/test_merge_collection.py
@@ -1,17 +1,22 @@
-from datetime import datetime, timezone, timedelta
+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.mergers import (
+from generalresearch.incite.mergers.base import (
MergeCollection,
MergeType,
)
-from test_utils.incite.conftest import mnt_filepath
-merge_types = list(e for e in MergeType if e != MergeType.TEST)
+if TYPE_CHECKING:
+ from generalresearch.incite.base import GRLDatasets
+
+merge_types = [e for e in MergeType if e != MergeType.TEST]
@pytest.mark.parametrize(
@@ -21,17 +26,20 @@ 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)],
)
),
)
class TestMergeCollection:
- def test_init(self, mnt_filepath, merge_type, offset, duration, start):
+ def test_init(
+ self,
+ merge_type: MergeType,
+ offset: str,
+ duration: timedelta,
+ start: datetime,
+ mnt_filepath: GRLDatasets,
+ ):
with pytest.raises(expected_exception=ValueError) as cm:
MergeCollection(archive_path=mnt_filepath.data_src)
assert "Must explicitly provide a merge_type" in str(cm.value)
@@ -42,7 +50,14 @@ class TestMergeCollection:
)
assert instance.merge_type == merge_type
- def test_items(self, mnt_filepath, merge_type, offset, duration, start):
+ def test_items(
+ self,
+ mnt_filepath: GRLDatasets,
+ merge_type: MergeType,
+ offset: str,
+ duration: timedelta,
+ start: datetime,
+ ):
instance = MergeCollection(
merge_type=merge_type,
offset=offset,
@@ -53,7 +68,14 @@ class TestMergeCollection:
assert len(instance.interval_range) == len(instance.items)
- def test_progress(self, mnt_filepath, merge_type, offset, duration, start):
+ def test_progress(
+ self,
+ mnt_filepath: GRLDatasets,
+ merge_type: MergeType,
+ offset: str,
+ duration: timedelta,
+ start: datetime,
+ ):
instance = MergeCollection(
merge_type=merge_type,
offset=offset,
@@ -67,7 +89,14 @@ class TestMergeCollection:
assert instance.progress.shape[1] == 7
assert instance.progress["group_by"].isnull().all()
- def test_schema(self, mnt_filepath, merge_type, offset, duration, start):
+ def test_schema(
+ self,
+ offset: str,
+ duration: timedelta,
+ start: datetime,
+ mnt_filepath: GRLDatasets,
+ merge_type: MergeType,
+ ):
instance = MergeCollection(
merge_type=merge_type,
archive_path=mnt_filepath.archive_path(enum_type=merge_type),
@@ -75,7 +104,14 @@ class TestMergeCollection:
assert isinstance(instance._schema, DataFrameSchema)
- def test_load(self, mnt_filepath, merge_type, offset, duration, start):
+ def test_load(
+ self,
+ mnt_filepath: GRLDatasets,
+ merge_type: MergeType,
+ offset: str,
+ duration: timedelta,
+ start: datetime,
+ ):
instance = MergeCollection(
merge_type=merge_type,
start=start,
@@ -87,7 +123,14 @@ class TestMergeCollection:
# Confirm that there are no archives available yet
assert instance.progress.has_archive.eq(False).all()
- def test_get_items(self, mnt_filepath, merge_type, offset, duration, start):
+ def test_get_items(
+ self,
+ mnt_filepath: GRLDatasets,
+ merge_type: MergeType,
+ offset: str,
+ duration: timedelta,
+ start: datetime,
+ ):
instance = MergeCollection(
start=start,
finished=start + duration,
diff --git a/tests/incite/mergers/test_merge_collection_item.py b/tests/incite/mergers/test_merge_collection_item.py
index 96f8789..baf1bc4 100644
--- a/tests/incite/mergers/test_merge_collection_item.py
+++ b/tests/incite/mergers/test_merge_collection_item.py
@@ -1,17 +1,19 @@
-from datetime import datetime, timezone, timedelta
+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 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
+from generalresearch.incite.mergers.base import MergeType
+
+if TYPE_CHECKING:
+ from generalresearch.incite.mergers.base import (
+ MergeCollection,
+ MergeCollectionItem,
+ )
@pytest.mark.parametrize(
@@ -26,7 +28,10 @@ from test_utils.incite.mergers.conftest import merge_collection
)
class TestMergeCollectionItem:
- def test_file_naming(self, merge_collection, offset, duration, start):
+ def test_file_naming(
+ self,
+ merge_collection: MergeCollection,
+ ):
assert len(merge_collection.items) == 25
items: list[MergeCollectionItem] = merge_collection.items
@@ -41,7 +46,10 @@ class TestMergeCollectionItem:
assert i._collection.offset in i.filename
assert i.start.strftime("%Y-%m-%d-%H-%M-%S") in i.filename
- def test_archives(self, merge_collection, offset, duration, start):
+ def test_archives(
+ self,
+ merge_collection: MergeCollection,
+ ):
assert len(merge_collection.items) == 25
for i in merge_collection.items:
@@ -51,10 +59,13 @@ class TestMergeCollectionItem:
assert not i.has_partial_archive()
assert i.has_archive() == i.path_exists(generic_path=i.path)
- res = set([i.should_archive() for i in merge_collection.items])
+ res = {i.should_archive() for i in merge_collection.items}
assert len(res) == 1
- def test_item_to_archive(self, merge_collection, offset, duration, start):
+ def test_item_to_archive(
+ self,
+ merge_collection: MergeCollection,
+ ):
for item in merge_collection.items:
item: MergeCollectionItem
assert not item.has_archive()
diff --git a/tests/incite/mergers/test_pop_ledger.py b/tests/incite/mergers/test_pop_ledger.py
index 6f96108..9ec188b 100644
--- a/tests/incite/mergers/test_pop_ledger.py
+++ b/tests/incite/mergers/test_pop_ledger.py
@@ -1,18 +1,28 @@
-from datetime import timedelta, datetime, timezone
+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 Optional
+from typing import TYPE_CHECKING
import pandas as pd
import pytest
-from distributed.utils_test import client_no_amm
+from dask.distributed import Client as DaskClient
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.mergers.conftest import pop_ledger_merge
-from test_utils.managers.ledger.conftest import create_main_accounts
+
+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(
@@ -27,25 +37,25 @@ 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(
self,
- client_no_amm,
- ledger_collection,
- pop_ledger_merge,
- product,
- user_factory,
- create_main_accounts,
- thl_lm,
- delete_df_collection,
- incite_item_factory,
- delete_ledger_db,
+ client_no_amm: DaskClient,
+ ledger_collection: LedgerDFCollection,
+ pop_ledger_merge: PopLedgerMerge,
+ product: Product,
+ user_factory: Callable[..., User],
+ create_main_accounts: Callable[..., None],
+ thl_ledger_manager: ThlLedgerManager,
+ delete_df_collection: Callable[..., None],
+ incite_item_factory: Callable[..., None],
+ delete_ledger_db: Callable[..., None],
):
from generalresearch.models.thl.ledger import LedgerAccount
@@ -79,19 +89,19 @@ class TestMergePOPLedger:
# --
- user_wallet_account: LedgerAccount = thl_lm.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()
)
- cash_account: LedgerAccount = thl_lm.get_account_cash()
- rev_account: LedgerAccount = thl_lm.get_account_task_complete_revenue()
item_finishes = [i.finish for i in ledger_collection.items]
item_finishes.sort(reverse=True)
last_item_finish = item_finishes[0]
# Pure SQL based lookups
- cash_balance: int = thl_lm.get_account_balance(account=cash_account)
- rev_balance: int = thl_lm.get_account_balance(account=rev_account)
+ cash_balance: int = thl_ledger_manager.get_account_balance(account=cash_account)
+ rev_balance: int = thl_ledger_manager.get_account_balance(account=rev_account)
assert cash_balance > rev_balance
# (1) Test Cash Account
@@ -129,39 +139,42 @@ class TestMergePOPLedger:
def test_pydantic_init(
self,
- client_no_amm,
- ledger_collection,
- pop_ledger_merge,
- mnt_filepath,
- product,
- user_factory,
- create_main_accounts,
- offset,
- duration,
- start,
- thl_lm,
- incite_item_factory,
- delete_df_collection,
- delete_ledger_db,
- session_collection,
+ client_no_amm: DaskClient,
+ ledger_collection: LedgerDFCollection,
+ pop_ledger_merge: PopLedgerMerge,
+ mnt_filepath: GRLDatasets,
+ product: Product,
+ user_factory: Callable[..., User],
+ create_main_accounts: Callable[..., None],
+ offset: str,
+ duration: timedelta,
+ start: datetime,
+ thl_ledger_manager: ThlLedgerManager,
+ incite_item_factory: Callable[..., None],
+ delete_df_collection: Callable[..., None],
+ delete_ledger_db: Callable[..., None],
+ session_collection: SessionDFCollection,
):
+ 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)
assert ledger_collection.finished is not None
assert isinstance(u.product, Product)
delete_ledger_db()
- create_main_accounts(),
+ create_main_accounts()
+
delete_df_collection(coll=ledger_collection)
- bp_account: LedgerAccount = thl_lm.get_account_or_create_bp_wallet(
+ bp_account: LedgerAccount = thl_ledger_manager.get_account_or_create_bp_wallet(
product=u.product
)
- cash_account: LedgerAccount = thl_lm.get_account_cash()
- rev_account: LedgerAccount = thl_lm.get_account_task_complete_revenue()
+ cash_account: LedgerAccount = thl_ledger_manager.get_account_cash()
+ rev_account: LedgerAccount = (
+ thl_ledger_manager.get_account_task_complete_revenue()
+ )
for item in ledger_collection.items:
incite_item_factory(item=item, user=u)
@@ -191,8 +204,10 @@ class TestMergePOPLedger:
assert instance.payout == instance.net == instance.bp_payment_credit
assert instance.available_balance < instance.net
assert instance.available_balance + instance.retainer == instance.net
- assert instance.balance == thl_lm.get_account_balance(bp_account)
- assert df["bp_payment.CREDIT"].sum() == thl_lm.get_account_balance(bp_account)
+ assert instance.balance == thl_ledger_manager.get_account_balance(bp_account)
+ assert df["bp_payment.CREDIT"].sum() == thl_ledger_manager.get_account_balance(
+ bp_account
+ )
# (2) Filter by the Cash Account
ddf = pop_ledger_merge.ddf(
@@ -205,7 +220,7 @@ class TestMergePOPLedger:
)
df: pd.DataFrame = client_no_amm.compute(collections=ddf, sync=True)
- cash_balance: int = thl_lm.get_account_balance(account=cash_account)
+ cash_balance: int = thl_ledger_manager.get_account_balance(account=cash_account)
assert df["bp_payment.CREDIT"].sum() == 0
assert cash_balance > 0
assert df["mp_payment.CREDIT"].sum() == 0
@@ -222,7 +237,7 @@ class TestMergePOPLedger:
)
df: pd.DataFrame = client_no_amm.compute(collections=ddf, sync=True)
- rev_balance: int = thl_lm.get_account_balance(account=rev_account)
+ rev_balance: int = thl_ledger_manager.get_account_balance(account=rev_account)
assert rev_balance == 0
assert df["bp_payment.CREDIT"].sum() == 0
assert df["mp_payment.DEBIT"].sum() == 0
@@ -230,27 +245,28 @@ class TestMergePOPLedger:
def test_resample(
self,
- client_no_amm,
- ledger_collection,
- pop_ledger_merge,
- mnt_filepath,
- user_factory,
- product,
- create_main_accounts,
- offset,
- duration,
- start,
- thl_lm,
- delete_df_collection,
- incite_item_factory,
+ client_no_amm: DaskClient,
+ ledger_collection: LedgerDFCollection,
+ pop_ledger_merge: PopLedgerMerge,
+ mnt_filepath: GRLDatasets,
+ user_factory: Callable[..., User],
+ product: Product,
+ create_main_accounts: Callable[..., None],
+ offset: str,
+ duration: timedelta,
+ start: datetime,
+ thl_ledger_manager: ThlLedgerManager,
+ delete_df_collection: Callable[..., None],
+ incite_item_factory: Callable[..., None],
):
- from generalresearch.models.thl.user import User
assert ledger_collection.finished is not None
delete_df_collection(coll=ledger_collection)
u1: User = user_factory(product=product)
- bp_account = thl_lm.get_account_or_create_bp_wallet(product=u1.product)
+ bp_account = thl_ledger_manager.get_account_or_create_bp_wallet(
+ product=u1.product
+ )
for item in ledger_collection.items:
incite_item_factory(user=u1, item=item)
@@ -280,7 +296,7 @@ class TestMergePOPLedger:
assert isinstance(df.index, pd.Index)
assert isinstance(df.index, pd.DatetimeIndex)
- bp_account_balance = thl_lm.get_account_balance(account=bp_account)
+ thl_ledger_manager.get_account_balance(account=bp_account)
# Initial sum
initial_sum = df.sum().sum()
diff --git a/tests/incite/mergers/test_ym_survey_merge.py b/tests/incite/mergers/test_ym_survey_merge.py
index 4c2df6b..d83a98c 100644
--- a/tests/incite/mergers/test_ym_survey_merge.py
+++ b/tests/incite/mergers/test_ym_survey_merge.py
@@ -1,25 +1,28 @@
-from datetime import timedelta, timezone, datetime
+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
+
+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
-from distributed.utils_test import (
- gen_cluster,
- client_no_amm,
- loop,
- loop_in_thread,
- cleanup,
- cluster_fixture,
- client,
-)
-
-from test_utils.incite.collections.conftest import wall_collection, session_collection
-from test_utils.incite.mergers.conftest import (
- enriched_session_merge,
- ym_survey_wall_merge,
-)
@pytest.mark.parametrize(
@@ -28,11 +31,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)],
)
),
)
@@ -46,18 +45,17 @@ class TestYMSurveyMerge:
def test_base(
self,
- client_no_amm,
- user_factory,
- product,
- ym_survey_wall_merge,
- wall_collection,
- session_collection,
- enriched_session_merge,
- delete_df_collection,
- incite_item_factory,
- thl_web_rr,
+ client_no_amm: DaskClient,
+ user_factory: Callable[..., User],
+ product: Product,
+ ym_survey_wall_merge: YMSurveyWallMerge,
+ wall_collection: WallDFCollection,
+ session_collection: SessionDFCollection,
+ enriched_session_merge: EnrichedSessionMerge,
+ delete_df_collection: Callable[..., None],
+ incite_item_factory: Callable[..., None],
+ thl_web_rr: PostgresConfig,
):
- from generalresearch.models.thl.user import User
delete_df_collection(coll=session_collection)
user: User = user_factory(product=product, created=session_collection.start)
@@ -85,10 +83,10 @@ class TestYMSurveyMerge:
assert enriched_session_merge.progress.has_archive.eq(True).all()
ddf = enriched_session_merge.ddf()
- df: pd.DataFrame = client_no_amm.compute(collections=ddf, sync=True)
+ df1: pd.DataFrame | None = client_no_amm.compute(collections=ddf, sync=True)
- assert isinstance(df, pd.DataFrame)
- assert not df.empty
+ assert isinstance(df1, pd.DataFrame)
+ assert not df1.empty
# --
@@ -102,18 +100,18 @@ class TestYMSurveyMerge:
# --
ddf = ym_survey_wall_merge.ddf()
- df: pd.DataFrame = client_no_amm.compute(collections=ddf, sync=True)
+ df2: pd.DataFrame | None = client_no_amm.compute(collections=ddf, sync=True)
- assert isinstance(df, pd.DataFrame)
- assert not df.empty
+ assert isinstance(df2, pd.DataFrame)
+ assert not df2.empty
# --
- assert df.product_id.nunique() == 1
- assert df.team_id.nunique() == 1
- assert df.source.nunique() > 1
+ assert df2.product_id.nunique() == 1
+ assert df2.team_id.nunique() == 1
+ assert df2.source.nunique() > 1
- started_min_ts = df.started.min()
- started_max_ts = df.started.max()
+ started_min_ts = df2.started.min()
+ started_max_ts = df2.started.max()
assert type(started_min_ts) is pd.Timestamp
assert type(started_max_ts) is pd.Timestamp
diff --git a/tests/incite/schemas/test_admin_responses.py b/tests/incite/schemas/test_admin_responses.py
index 43aa399..d2658ea 100644
--- a/tests/incite/schemas/test_admin_responses.py
+++ b/tests/incite/schemas/test_admin_responses.py
@@ -1,15 +1,17 @@
-from datetime import datetime, timezone, timedelta
+from __future__ import annotations
+
+from datetime import UTC, datetime, timedelta
from random import sample
-from typing import List
import numpy as np
import pandas as pd
+import pandera as pa
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
@@ -17,12 +19,14 @@ from generalresearch.locales import Localelator
class TestAdminPOPSchema:
schema_df = empty_dataframe_from_schema(AdminPOPSchema)
countries = list(Localelator().get_all_countries())[:5]
- dates = [datetime(year=2024, month=1, day=i, tzinfo=None) for i in range(1, 10)]
+ dates = [
+ datetime(year=2024, month=1, day=i, tzinfo=None) for i in range(1, 10) # noqa
+ ]
@classmethod
def assign_valid_vals(cls, df: pd.DataFrame) -> pd.DataFrame:
for c in df.columns:
- check_attrs: dict = AdminPOPSchema.columns[c].checks[0].statistics
+ check_attrs = AdminPOPSchema.columns[c].checks[0].statistics
df[c] = np.random.randint(
check_attrs["min_value"], check_attrs["max_value"], df.shape[0]
)
@@ -30,7 +34,7 @@ class TestAdminPOPSchema:
return df
def test_empty(self):
- with pytest.raises(Exception):
+ with pytest.raises(pa.errors.SchemaError):
AdminPOPSchema.validate(pd.DataFrame())
def test_new_empty_df(self):
@@ -43,7 +47,7 @@ class TestAdminPOPSchema:
def test_valid(self):
# (1) Works with raw naive datetime
dates = [
- datetime(year=2024, month=1, day=i, tzinfo=None).isoformat()
+ datetime(year=2024, month=1, day=i, tzinfo=None).isoformat() # noqa
for i in range(1, 10)
]
df = pd.DataFrame(
@@ -58,7 +62,10 @@ class TestAdminPOPSchema:
assert isinstance(df, pd.DataFrame)
# (2) Works with isoformat naive datetime
- dates = [datetime(year=2024, month=1, day=i, tzinfo=None) for i in range(1, 10)]
+ dates = [
+ datetime(year=2024, month=1, day=i, tzinfo=None) # noqa
+ for i in range(1, 10)
+ ]
df = pd.DataFrame(
index=pd.MultiIndex.from_product(
iterables=[dates, self.countries], names=["index0", "index1"]
@@ -72,8 +79,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 +91,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)]
- assert all([ts.tz is None for ts in timestmaps])
+ 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(
@@ -125,12 +131,12 @@ class TestAdminPOPSchema:
df = self.assign_valid_vals(df)
vals = [i for i in df.index.get_level_values(1)]
- assert all([isinstance(v, float) for v in vals])
+ assert all(isinstance(v, float) for v in vals)
df = AdminPOPSchema.validate(df, lazy=True)
vals = [i for i in df.index.get_level_values(1)]
- assert all([isinstance(v, str) for v in vals])
+ assert all(isinstance(v, str) for v in vals)
# --- int to str ---
@@ -144,12 +150,12 @@ class TestAdminPOPSchema:
df = self.assign_valid_vals(df)
vals = [i for i in df.index.get_level_values(1)]
- assert all([isinstance(v, int) for v in vals])
+ assert all(isinstance(v, int) for v in vals)
df = AdminPOPSchema.validate(df, lazy=True)
vals = [i for i in df.index.get_level_values(1)]
- assert all([isinstance(v, str) for v in vals])
+ assert all(isinstance(v, str) for v in vals)
# a = 1
assert isinstance(df, pd.DataFrame)
@@ -157,9 +163,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 +177,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)]
- assert all([ts.tz is None for ts in timestmaps])
+ 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,13 +194,13 @@ 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)]
- assert all([ts.tz is None for ts in timestmaps])
+ 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):
df = pd.DataFrame(
diff --git a/tests/incite/schemas/test_thl_web.py b/tests/incite/schemas/test_thl_web.py
index 7f4434b..9b34ce0 100644
--- a/tests/incite/schemas/test_thl_web.py
+++ b/tests/incite/schemas/test_thl_web.py
@@ -16,7 +16,7 @@ class TestWallSchema:
df = pd.DataFrame(columns=THLWallSchema.columns.keys())
- with pytest.raises(SchemaError) as cm:
+ with pytest.raises(SchemaError):
THLWallSchema.validate(df)
def test_no_rows(self):
@@ -24,7 +24,7 @@ class TestWallSchema:
df = pd.DataFrame(index=["uuid"], columns=THLWallSchema.columns.keys())
- with pytest.raises(SchemaError) as cm:
+ with pytest.raises(SchemaError):
THLWallSchema.validate(df)
def test_new_empty_df(self):
@@ -50,7 +50,7 @@ class TestSessionSchema:
df = pd.DataFrame(columns=THLSessionSchema.columns.keys())
df.set_index("uuid", inplace=True)
- with pytest.raises(SchemaError) as cm:
+ with pytest.raises(SchemaError):
THLSessionSchema.validate(df)
def test_no_rows(self):
@@ -58,7 +58,7 @@ class TestSessionSchema:
df = pd.DataFrame(index=["id"], columns=THLSessionSchema.columns.keys())
- with pytest.raises(SchemaError) as cm:
+ with pytest.raises(SchemaError):
THLSessionSchema.validate(df)
def test_new_empty_df(self):
diff --git a/tests/incite/test_collection_base.py b/tests/incite/test_collection_base.py
index 7e6605f..1a664a2 100644
--- a/tests/incite/test_collection_base.py
+++ b/tests/incite/test_collection_base.py
@@ -1,7 +1,10 @@
-from datetime import datetime, timedelta, timezone
+from __future__ import annotations
+
+from datetime import UTC, datetime, timedelta, timezone
from os.path import exists as pexists
from os.path import join as pjoin
from pathlib import Path
+from typing import TYPE_CHECKING
from uuid import uuid4
import numpy as np
@@ -10,21 +13,21 @@ 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=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)
+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)
+AGO_2HR = (datetime.now(tz=UTC) - timedelta(hours=2)).replace(microsecond=0)
class TestCollectionBase:
- def test_init(self, mnt_filepath):
+ def test_init(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
assert instance.df.empty is True
- def test_init_df(self, mnt_filepath):
+ def test_init_df(self, mnt_filepath: GRLDatasets):
# Only an empty pd.DataFrame can ever be provided
instance = CollectionBase(
df=pd.DataFrame({}), archive_path=mnt_filepath.data_src
@@ -46,11 +49,11 @@ class TestCollectionBase:
)
assert "Do not provide a pd.DataFrame" in str(cm.value)
- def test_init_start(self, mnt_filepath):
+ def test_init_start(self, mnt_filepath: GRLDatasets):
with pytest.raises(expected_exception=ValueError) as cm:
cm: ExceptionInfo
CollectionBase(
- 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 +69,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
@@ -79,7 +80,7 @@ class TestCollectionBase:
cm.value
)
- def test_init_archive_path(self, mnt_filepath):
+ def test_init_archive_path(self, mnt_filepath: GRLDatasets):
"""DirectoryPath is apparently smart enough to confirm that the
directory path exists.
"""
@@ -104,7 +105,7 @@ class TestCollectionBase:
CollectionBase(archive_path=new_path)
assert "Path does not point to a directory" in str(cm.value)
- def test_init_offset(self, mnt_filepath):
+ def test_init_offset(self, mnt_filepath: GRLDatasets):
with pytest.raises(expected_exception=ValueError) as cm:
cm: ExceptionInfo
CollectionBase(offset="1:X", archive_path=mnt_filepath.data_src)
@@ -112,7 +113,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:
@@ -123,14 +124,14 @@ class TestCollectionBase:
class TestCollectionBaseProperties:
- def test_items(self, mnt_filepath):
+ def test_items(self, mnt_filepath: GRLDatasets):
with pytest.raises(expected_exception=NotImplementedError) as cm:
cm: ExceptionInfo
instance = CollectionBase(archive_path=mnt_filepath.data_src)
- x = instance.items
+ _ = instance.items
assert "Must override" in str(cm.value)
- def test_interval_range(self, mnt_filepath):
+ def test_interval_range(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
# Private method requires the end parameter
with pytest.raises(expected_exception=AssertionError) as cm:
@@ -145,14 +146,14 @@ 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
assert res.is_monotonic_increasing
assert res.is_unique
- def test_interval_range2(self, mnt_filepath):
+ def test_interval_range2(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
assert isinstance(instance.interval_range, list)
@@ -171,16 +172,16 @@ class TestCollectionBaseProperties:
)
assert len(instance.interval_range) == 2
- def test_progress(self, mnt_filepath):
+ def test_progress(self, mnt_filepath: GRLDatasets):
with pytest.raises(expected_exception=NotImplementedError) as cm:
cm: ExceptionInfo
instance = CollectionBase(
start=AGO_15min, offset="3min", archive_path=mnt_filepath.data_src
)
- x = instance.progress
+ _ = instance.progress
assert "Must override" in str(cm.value)
- def test_progress2(self, mnt_filepath):
+ def test_progress2(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(
start=AGO_2HR,
offset="15min",
@@ -189,10 +190,10 @@ class TestCollectionBaseProperties:
assert instance.df.empty
with pytest.raises(expected_exception=NotImplementedError) as cm:
- df = instance.progress
+ _ = instance.progress
assert "Must override" in str(cm.value)
- def test_items2(self, mnt_filepath):
+ def test_items2(self, mnt_filepath: GRLDatasets):
"""There can't be a test for this because the Items need a path whic
isn't possible in the generic form
"""
@@ -202,7 +203,7 @@ class TestCollectionBaseProperties:
with pytest.raises(expected_exception=NotImplementedError) as cm:
cm: ExceptionInfo
- items = instance.items
+ _ = instance.items
assert "Must override" in str(cm.value)
# item = items[-3]
@@ -213,19 +214,19 @@ class TestCollectionBaseProperties:
# assert str(df.product_id.dtype) == "object"
# assert str(ddf.product_id.dtype) == "string"
- def test_items3(self, mnt_filepath):
+ def test_items3(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(
start=AGO_2HR,
offset="15min",
archive_path=mnt_filepath.data_src,
)
with pytest.raises(expected_exception=NotImplementedError) as cm:
- item = instance.items[0]
+ _ = instance.items[0]
assert "Must override" in str(cm.value)
class TestCollectionBaseMethodsCleanup:
- def test_fetch_force_rr_latest(self, mnt_filepath):
+ def test_fetch_force_rr_latest(self, mnt_filepath: GRLDatasets):
coll = CollectionBase(archive_path=mnt_filepath.data_src)
with pytest.raises(expected_exception=Exception) as cm:
@@ -233,7 +234,7 @@ class TestCollectionBaseMethodsCleanup:
coll.fetch_force_rr_latest(sources=[])
assert "Must override" in str(cm.value)
- def test_fetch_all_paths(self, mnt_filepath):
+ def test_fetch_all_paths(self, mnt_filepath: GRLDatasets):
coll = CollectionBase(archive_path=mnt_filepath.data_src)
with pytest.raises(expected_exception=NotImplementedError) as cm:
@@ -244,19 +245,19 @@ class TestCollectionBaseMethodsCleanup:
assert "Must override" in str(cm.value)
-class TestCollectionBaseMethodsCleanup:
+class TestCollectionBaseMethodsCleanup2:
@pytest.mark.skip
- def test_cleanup_partials(self, mnt_filepath):
+ def test_cleanup_partials(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
assert instance.cleanup_partials() is None # it doesn't return anything
- def test_clear_tmp_archives(self, mnt_filepath):
+ def test_clear_tmp_archives(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
assert instance.clear_tmp_archives() is None # it doesn't return anything
@pytest.mark.skip
- def test_clear_corrupt_archives(self, mnt_filepath):
+ def test_clear_corrupt_archives(self, mnt_filepath: GRLDatasets):
"""TODO: expand this so it actually has corrupt archives that we
check to see if they're removed
"""
@@ -264,14 +265,14 @@ class TestCollectionBaseMethodsCleanup:
assert instance.clear_corrupt_archives() is None # it doesn't return anything
@pytest.mark.skip
- def test_rebuild_symlinks(self, mnt_filepath):
+ def test_rebuild_symlinks(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
assert instance.rebuild_symlinks() is None
class TestCollectionBaseMethodsSourceTiming:
- def test_get_item(self, mnt_filepath):
+ def test_get_item(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
i = pd.Interval(left=1, right=2, closed="left")
@@ -279,40 +280,40 @@ class TestCollectionBaseMethodsSourceTiming:
instance.get_item(interval=i)
assert "Must override" in str(cm.value)
- def test_get_item_start(self, mnt_filepath):
+ def test_get_item_start(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
- dt = datetime.now(tz=timezone.utc)
+ dt = datetime.now(tz=UTC)
start = pd.Timestamp(dt)
with pytest.raises(expected_exception=NotImplementedError) as cm:
instance.get_item_start(start=start)
assert "Must override" in str(cm.value)
- def test_get_items(self, mnt_filepath):
+ def test_get_items(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
- dt = datetime.now(tz=timezone.utc)
+ dt = datetime.now(tz=UTC)
with pytest.raises(expected_exception=NotImplementedError) as cm:
instance.get_items(since=dt)
assert "Must override" in str(cm.value)
- def test_get_items_from_year(self, mnt_filepath):
+ def test_get_items_from_year(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
with pytest.raises(expected_exception=NotImplementedError) as cm:
instance.get_items_from_year(year=2020)
assert "Must override" in str(cm.value)
- def test_get_items_last90(self, mnt_filepath):
+ def test_get_items_last90(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
with pytest.raises(expected_exception=NotImplementedError) as cm:
instance.get_items_last90()
assert "Must override" in str(cm.value)
- def test_get_items_last365(self, mnt_filepath):
+ def test_get_items_last365(self, mnt_filepath: GRLDatasets):
instance = CollectionBase(archive_path=mnt_filepath.data_src)
with pytest.raises(expected_exception=NotImplementedError) as cm:
diff --git a/tests/incite/test_collection_base_item.py b/tests/incite/test_collection_base_item.py
index e5d1d02..b9f1c26 100644
--- a/tests/incite/test_collection_base_item.py
+++ b/tests/incite/test_collection_base_item.py
@@ -1,6 +1,9 @@
-from datetime import datetime, timezone
+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,10 +13,13 @@ from pydantic import ValidationError
from generalresearch.incite.base import CollectionItemBase
+if TYPE_CHECKING:
+ from generalresearch.incite.base import GRLDatasets
+
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 +31,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)
@@ -40,20 +46,20 @@ class TestCollectionItemBaseProperties:
def test_finish(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.finish
+ with pytest.raises(expected_exception=AttributeError):
+ _ = instance.finish
def test_interval(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.interval
+ with pytest.raises(expected_exception=AttributeError):
+ _ = instance.interval
def test_filename(self):
instance = CollectionItemBase()
with pytest.raises(expected_exception=NotImplementedError) as cm:
- res = instance.filename
+ _ = instance.filename
assert "Do not use CollectionItemBase directly" in str(cm.value)
@@ -61,7 +67,7 @@ class TestCollectionItemBaseProperties:
instance = CollectionItemBase()
with pytest.raises(expected_exception=NotImplementedError) as cm:
- res = instance.filename
+ _ = instance.filename
assert "Do not use CollectionItemBase directly" in str(cm.value)
@@ -69,27 +75,27 @@ class TestCollectionItemBaseProperties:
instance = CollectionItemBase()
with pytest.raises(expected_exception=NotImplementedError) as cm:
- res = instance.filename
+ _ = instance.filename
assert "Do not use CollectionItemBase directly" in str(cm.value)
def test_path(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.path
+ with pytest.raises(expected_exception=AttributeError):
+ _ = instance.path
def test_partial_path(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.partial_path
+ with pytest.raises(expected_exception=AttributeError):
+ _ = instance.partial_path
def test_empty_path(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.empty_path
+ with pytest.raises(expected_exception=AttributeError):
+ _ = instance.empty_path
class TestCollectionItemBaseMethods:
@@ -106,41 +112,41 @@ class TestCollectionItemBaseMethods:
instance = CollectionItemBase()
with pytest.raises(expected_exception=NotImplementedError) as cm:
- res = instance.tmp_filename()
+ _ = instance.tmp_filename()
assert "Do not use CollectionItemBase directly" in str(cm.value)
def test_tmp_path(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.tmp_path()
+ with pytest.raises(expected_exception=AttributeError):
+ instance.tmp_path()
def test_is_empty(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.is_empty()
+ with pytest.raises(expected_exception=AttributeError):
+ instance.is_empty()
def test_has_empty(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.has_empty()
+ with pytest.raises(expected_exception=AttributeError):
+ instance.has_empty()
def test_has_partial_archive(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.has_partial_archive()
+ with pytest.raises(expected_exception=AttributeError):
+ instance.has_partial_archive()
@pytest.mark.parametrize("include_empty", [True, False])
- def test_has_archive(self, include_empty):
+ def test_has_archive(self, include_empty: bool):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.has_archive(include_empty=include_empty)
+ with pytest.raises(expected_exception=AttributeError):
+ instance.has_archive(include_empty=include_empty)
- def test_delete_archive_file(self, mnt_filepath):
+ def test_delete_archive_file(self, mnt_filepath: GRLDatasets):
path1 = Path(pjoin(mnt_filepath.data_src, f"{uuid4().hex}.zip"))
# Confirm it doesn't exist, and that delete_archive() doesn't throw
@@ -155,7 +161,7 @@ class TestCollectionItemBaseMethods:
CollectionItemBase.delete_archive(generic_path=path1)
assert not path1.exists()
- def test_delete_archive_dir(self, mnt_filepath):
+ def test_delete_archive_dir(self, mnt_filepath: GRLDatasets):
path1 = Path(pjoin(mnt_filepath.data_src, f"{uuid4().hex}"))
# Confirm it doesn't exist, and that delete_archive() doesn't throw
@@ -174,20 +180,20 @@ class TestCollectionItemBaseMethods:
def test_should_archive(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.should_archive()
+ with pytest.raises(expected_exception=AttributeError):
+ _ = instance.should_archive()
def test_set_empty(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.set_empty()
+ with pytest.raises(expected_exception=AttributeError):
+ _ = instance.set_empty()
def test_valid_archive(self):
instance = CollectionItemBase()
- with pytest.raises(expected_exception=AttributeError) as cm:
- res = instance.valid_archive(generic_path=None, sample=None)
+ with pytest.raises(expected_exception=AttributeError):
+ _ = instance.valid_archive(generic_path=None, sample=None)
class TestCollectionItemBaseMethodsORM:
@@ -197,11 +203,11 @@ class TestCollectionItemBaseMethodsORM:
pass
@pytest.mark.parametrize("is_partial", [True, False])
- def test_to_archive(self, is_partial):
+ def test_to_archive(self, is_partial: bool):
instance = CollectionItemBase()
with pytest.raises(expected_exception=NotImplementedError) as cm:
- res = instance.to_archive(
+ _ = instance.to_archive(
ddf=dd.from_pandas(data=pd.DataFrame()), is_partial=is_partial
)
assert "Must override" in str(cm.value)
diff --git a/tests/incite/test_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..04d0bb2 100644
--- a/tests/incite/test_interval_idx.py
+++ b/tests/incite/test_interval_idx.py
@@ -1,12 +1,13 @@
+from datetime import UTC, datetime
+
import pandas as pd
-from datetime import datetime, timezone, timedelta
class TestIntervalIndex:
def test_init(self):
- start = datetime(year=2000, month=1, day=1)
- end = datetime(year=2000, month=1, day=10)
+ start = datetime(year=2000, month=1, day=1, tzinfo=UTC)
+ end = datetime(year=2000, month=1, day=10, tzinfo=UTC)
iv_r: pd.IntervalIndex = pd.interval_range(
start=start, end=end, freq="1d", closed="left"
@@ -17,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_authentication.py b/tests/managers/gr/test_authentication.py
index 53b6931..1310c79 100644
--- a/tests/managers/gr/test_authentication.py
+++ b/tests/managers/gr/test_authentication.py
@@ -1,120 +1,150 @@
import logging
-from random import randint
+from collections.abc import Callable
from uuid import uuid4
import pytest
-from generalresearch.models.gr.authentication import GRUser
-from test_utils.models.conftest import gr_user
+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):
- from generalresearch.models.gr.authentication import GRUser
-
- 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_token_manager: GRTokenManager):
+ assert gr_token_manager.create(user_id=gr_user.id) is None
- token = gr_tm.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, gr_tm):
- assert gr_tm.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_tm.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, gr_tm, gr_db, gr_redis_config):
- from generalresearch.models.gr.authentication import GRToken
+ def test_prefetch_user(
+ self,
+ gr_user: GRUser,
+ gr_token_manager: GRTokenManager,
+ gr_db: PostgresConfig,
+ gr_redis_config: RedisConfig,
+ ):
- gr_tm.create(user_id=gr_user.id)
+ gr_token_manager.create(user_id=gr_user.id)
- token: GRToken = gr_tm.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)
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_token_manager: GRTokenManager,
+ ):
+ 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_tm.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_tm.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...")
- 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="...",
diff --git a/tests/managers/gr/test_business.py b/tests/managers/gr/test_business.py
index 7eb77f8..022086a 100644
--- a/tests/managers/gr/test_business.py
+++ b/tests/managers/gr/test_business.py
@@ -1,30 +1,52 @@
+from collections.abc import Callable
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
-from test_utils.models.conftest import business
+from generalresearch.models.gr.business import (
+ Business,
+ BusinessAddress,
+ BusinessBankAccount,
+)
+from generalresearch.models.gr.definitions import TransferMethod
+from generalresearch.models.gr.team import Team
+
+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:
- def test_init(self, business_bank_account_manager, gr_db):
- assert business_bank_account_manager.pg_config == gr_db
+ def test_init(
+ self,
+ gr_business_bank_account_manager: BusinessBankAccountManager,
+ gr_db: PostgresConfig,
+ ):
+ assert gr_business_bank_account_manager.pg_config == gr_db
- def test_create(self, business, business_bank_account_manager):
- from generalresearch.models.gr.business import (
- TransferMethod,
- BusinessBankAccount,
- )
+ def test_create(
+ self,
+ gr_business: Business,
+ gr_business_bank_account_manager: BusinessBankAccountManager,
+ ):
- instance = business_bank_account_manager.create(
- business_id=business.id,
+ instance = gr_business_bank_account_manager.create(
+ business_id=gr_business.id,
uuid=uuid4().hex,
transfer_method=TransferMethod.ACH,
)
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)
@@ -35,42 +57,49 @@ class TestBusinessBankAccountManager:
class TestBusinessAddressManager:
- def test_create(self, business, business_address_manager):
- from generalresearch.models.gr.business import BusinessAddress
-
- res = business_address_manager.create(uuid=uuid4().hex, business_id=business.id)
+ def test_create(
+ self, gr_business: Business, gr_business_address_manager: BusinessAddressManager
+ ):
+ assert gr_business.id
+ res = gr_business_address_manager.create(
+ uuid=uuid4().hex, business_id=gr_business.id
+ )
assert isinstance(res, BusinessAddress)
assert isinstance(res.id, int)
class TestBusinessManager:
- def test_create(self, business_manager):
- from generalresearch.models.gr.business import Business
+ 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)
- def test_get_or_create(self, business_manager):
+ 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):
- res1 = business_manager.get_all()
+ def test_get_all(
+ self,
+ gr_business_manager: BusinessManager,
+ gr_business_factory: Callable[..., Business],
+ ):
+ res1 = gr_business_manager.get_all()
assert isinstance(res1, list)
- business_manager.create_dummy()
- res2 = business_manager.get_all()
+ gr_business_factory()
+ res2 = gr_business_manager.get_all()
assert len(res1) == len(res2) - 1
@pytest.mark.skip(reason="TODO")
@@ -78,53 +107,65 @@ class TestBusinessManager:
pass
def test_get_by_user_id(
- self, business_manager, gr_user, team_manager, membership_manager
+ self,
+ gr_business_manager: BusinessManager,
+ gr_user: GRUser,
+ gr_team_manager: TeamManager,
+ gr_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, but don't add it to anything
- b1 = business_manager.create_dummy()
- res = business_manager.get_by_user_id(user_id=gr_user.id)
+ # Create a business: Business, but don't add it to anything
+ b1 = gr_business_factory()
+ 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 = team_manager.create_dummy()
- res = business_manager.get_by_user_id(user_id=gr_user.id)
+ t1 = gr_team_factory()
+ 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
- m1 = membership_manager.create(team=t1, gr_user=gr_user)
- res = business_manager.get_by_user_id(user_id=gr_user.id)
+ _ = 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
# 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)
+ gr_team_manager.add_business(team=t1, business=b1)
+ res = gr_business_manager.get_by_user_id(user_id=gr_user.id)
assert len(res) == 1
# Add another Business to the Team!
- b2 = business_manager.create_dummy()
- team_manager.add_business(team=t1, business=b2)
- res = business_manager.get_by_user_id(user_id=gr_user.id)
+ b2 = gr_business_factory()
+ gr_team_manager.add_business(team=t1, business=b2)
+ res = gr_business_manager.get_by_user_id(user_id=gr_user.id)
assert len(res) == 2
@pytest.mark.skip(reason="TODO")
def test_get_uuids_by_user_id(self):
pass
- def test_get_by_uuid(self, business, business_manager):
- instance = business_manager.get_by_uuid(business_uuid=business.uuid)
- assert business.id == instance.id
+ def test_get_by_uuid(
+ self, gr_business: Business, gr_business_manager: BusinessManager
+ ):
+ 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, business, business_manager):
- instance = business_manager.get_by_id(business_id=business.id)
- assert business.uuid == instance.uuid
+ 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
- def test_cache_key(self, 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
@@ -133,7 +174,7 @@ class TestBusinessManager:
# business = BusinessManager.create(
# uuid=b_uuid,
# name=f"test-{b_uuid[:6]}")
- # assert isinstance(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/gr/test_team.py b/tests/managers/gr/test_team.py
index 9215da4..878a9ca 100644
--- a/tests/managers/gr/test_team.py
+++ b/tests/managers/gr/test_team.py
@@ -1,105 +1,135 @@
+from __future__ import annotations
+
+from collections.abc import Callable
+from typing import TYPE_CHECKING
from uuid import uuid4
-from test_utils.models.conftest import team
+from generalresearch.models.gr.authentication import GRUser
+from generalresearch.models.gr.team import Membership, Team
+
+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:
- def test_init(self, membership_manager, gr_db):
- 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:
- def test_init(self, team_manager, gr_db):
- 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):
+ 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):
- 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):
- from generalresearch.models.gr.team import Team
+ 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_manager, gr_um, gr_db, gr_redis_config):
- from generalresearch.models.gr.authentication import GRUser
- from generalresearch.models.gr.team import Membership
+ def test_add_user(
+ self,
+ 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):
- from generalresearch.models.gr.team import Team
+ 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):
- from generalresearch.models.gr.team import Team
+ 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_manager, gr_um):
- from generalresearch.models.gr.authentication import GRUser
- from generalresearch.models.gr.team import Team
-
- user: GRUser = gr_um.create_dummy()
- team_manager.add_user(team=team, gr_user=user)
+ def test_get_by_user(
+ 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
+ )
- 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_token,
- gr_user,
- membership,
- product_factory,
- membership_factory,
- team,
- thl_web_rr,
- gr_redis_config,
- gr_db,
+ gr_user: GRUser,
+ product_factory: Callable[..., Product],
+ 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,
diff --git a/tests/managers/leaderboard.py b/tests/managers/leaderboard.py
index 4d32dd0..fad0b6b 100644
--- a/tests/managers/leaderboard.py
+++ b/tests/managers/leaderboard.py
@@ -1,8 +1,12 @@
+from __future__ import annotations
+
import os
import time
import zoneinfo
-from datetime import datetime, timezone
+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
@@ -10,9 +14,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 +23,13 @@ from generalresearch.models.thl.product import (
PayoutConfig,
PayoutTransformation,
PayoutTransformationPercentArgs,
+ Product,
)
+from generalresearch.models.thl.session import Session
+from generalresearch.models.thl.user import User
+
+if TYPE_CHECKING:
+ from generalresearch.redis_helper import RedisConfig
# random uuid for leaderboard tests
product_id = uuid4().hex
@@ -44,7 +51,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,
@@ -63,7 +72,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,
@@ -74,59 +83,68 @@ 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_config: RedisConfig) -> Callable[..., None]:
+ thl_redis = thl_redis_config.create_redis_client()
+
+ 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: Callable[..., None], thl_redis_config: RedisConfig
+ ):
+ thl_redis = thl_redis_config.create_redis_client()
+ setup_leaderboards()
- def test_leaderboard_manager(self, setup_leaderboards, thl_redis):
country_iso = "us"
board_code = LeaderboardCode.COMPLETE_COUNT
freq = LeaderboardFrequency.DAILY
@@ -136,6 +154,7 @@ class TestLeaderboards:
freq=freq,
product_id=product_id,
country_iso=country_iso,
+ # This is supposed to not have a timezone. @max don't change it
within_time=datetime(2025, 2, 5, 0, 0, 0),
)
lb = m.get_leaderboard()
@@ -152,7 +171,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),
@@ -164,7 +183,12 @@ 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_config: RedisConfig
+ ):
+ thl_redis = thl_redis_config.create_redis_client()
+ setup_leaderboards()
+
country_iso = "us"
board_code = LeaderboardCode.COMPLETE_COUNT
freq = LeaderboardFrequency.DAILY
@@ -174,7 +198,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 +215,15 @@ 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_config: RedisConfig,
+ ):
+ thl_redis = thl_redis_config.create_redis_client()
+ setup_leaderboards()
+
hit_leaderboards(redis_client=thl_redis, session=session_factory())
for freq in [
@@ -205,7 +237,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 +248,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 +259,21 @@ 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_config: RedisConfig,
):
+ thl_redis = thl_redis_config.create_redis_client()
+ setup_leaderboards()
+
session = session_factory(product_user_id="zzz")
hit_leaderboards(redis_client=thl_redis, session=session)
m = LeaderboardManager(
@@ -244,24 +282,21 @@ 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_config: RedisConfig):
+ thl_redis = thl_redis_config.create_redis_client()
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
@@ -270,5 +305,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/network/__init__.py b/tests/managers/network/__init__.py
deleted file mode 100644
index e69de29..0000000
--- a/tests/managers/network/__init__.py
+++ /dev/null
diff --git a/tests/managers/network/test_label.py b/tests/managers/network/test_label.py
deleted file mode 100644
index 5b9a790..0000000
--- a/tests/managers/network/test_label.py
+++ /dev/null
@@ -1,202 +0,0 @@
-import ipaddress
-
-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,
- IPLabelSource,
- IPLabelMetadata,
-)
-from generalresearch.models.thl.ipinfo import normalize_ip
-
-fake = faker.Faker()
-
-
-@pytest.fixture
-def ip_label(utc_now) -> 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):
- 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):
- # 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
-):
- # 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, prefix = 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/managers/test_events.py b/tests/managers/test_events.py
index a0fab38..5ebd015 100644
--- a/tests/managers/test_events.py
+++ b/tests/managers/test_events.py
@@ -1,60 +1,48 @@
-import random
+from __future__ import annotations
+
+import math
import time
-from datetime import timedelta, datetime, timezone
+from collections.abc import Callable
+from datetime import UTC, datetime, timedelta
from decimal import Decimal
-from functools import partial
-from typing import Optional
+from typing import TYPE_CHECKING
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.definitions import Source
from generalresearch.models.events import (
- MessageKind,
- EventType,
AggregateBySource,
+ EventType,
MaxGaugeBySource,
+ MessageKind,
)
from generalresearch.models.legacy.bucket import Bucket
+from generalresearch.models.thl import Product
from generalresearch.models.thl.definitions import Status, StatusCode1
from generalresearch.models.thl.session import Session, Wall
from generalresearch.models.thl.user import User
+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
@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):
- return partial(create_dummy, product_id=product_id)
-
-
-@pytest.fixture(scope="function")
-def event_subscriber(thl_redis_config, 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(
- product_id: Optional[str] = None, product_user_id: Optional[str] = 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),
- user_id=random.randint(0, floor(2**32 / 2)),
- )
-
-
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,
@@ -63,7 +51,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_factory,
+ user_factory: Callable[..., User],
+ ):
event_manager.clear_global_user_stats()
user1: User = user_factory()
@@ -72,7 +65,7 @@ class TestActiveUsers:
event_manager.handle_user(user1)
event_manager.handle_user(user1)
- res = event_manager.get_user_stats(product_id)
+ res = event_manager.get_user_stats(user1.product_id)
assert res == {
"active_users_last_1h": 1,
"active_users_last_24h": 1,
@@ -88,21 +81,23 @@ class TestActiveUsers:
}
# Create a 2nd user in another product
- product_id2 = uuid4().hex
- user2: User = user_factory(product_id=product_id2)
+ product2 = product_factory()
+ user2: User = user_factory(product=product2)
+ 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)
# And now each have 1 active user
- assert event_manager.get_user_stats(product_id) == {
+ assert event_manager.get_user_stats(user1.product_id) == {
"active_users_last_1h": 1,
"active_users_last_24h": 1,
"signups_last_24h": 1,
"in_progress_users": 0,
}
# user2 was created older than 24hrs ago
- assert event_manager.get_user_stats(product_id2) == {
+ assert event_manager.get_user_stats(user2.product_id) == {
"active_users_last_1h": 1,
"active_users_last_24h": 1,
"signups_last_24h": 0,
@@ -116,10 +111,16 @@ class TestActiveUsers:
"in_progress_users": 0,
}
- def test_inprogress(self, event_manager, product_id, user_factory):
+ def test_inprogress(
+ self,
+ event_manager: EventManager,
+ user_factory: Callable[..., User],
+ product
+ ):
event_manager.clear_global_user_stats()
- user1: User = user_factory()
- user2: User = user_factory()
+ user1: User = user_factory(product=product)
+ user2: User = user_factory(product=product)
+ product_id = product.id
# No matter how many times we do this, they're only active once
event_manager.mark_user_inprogress(user1)
@@ -139,9 +140,14 @@ 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,
+ user_factory: Callable[..., User],
+ ):
event_manager.clear_global_user_stats()
user1: User = user_factory()
+ product_id = user1.product_id
event_manager.handle_user(user1)
event_manager.mark_user_inprogress(user1)
sec_24hr = timedelta(hours=24).total_seconds()
@@ -166,8 +172,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,
@@ -186,10 +191,18 @@ 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: EventManager,
+ product_factory: Callable[..., Product],
+ user_factory: Callable[..., User],
+ utc_now: datetime,
+ utc_hour_ago: datetime,
+ ):
event_manager.clear_global_session_stats()
-
- user: User = user_factory()
+ product = product_factory()
+ product_id = product.id
+ user: User = user_factory(product=product)
session = Session(
country_iso="us",
started=utc_hour_ago + timedelta(minutes=10),
@@ -266,29 +279,29 @@ class TestSessionStats:
field_name = str(field)
assert res == {field_name: "1"}
assert (
- 3600 - 60 < event_manager.redis_client.httl(name, field_name)[0] < 3600 + 60
+ 3600 - 61 < event_manager.redis_client.httl(name, field_name)[0] < 3600 + 60
)
# Second BP, fail
- product_id2 = uuid4().hex
- user2: User = user_factory(product_id=product_id2)
+ product2 = product_factory()
+ user2: User = user_factory(product=product2)
session3 = Session(
country_iso="us",
started=utc_now - timedelta(minutes=1),
user=user2,
)
- event_manager.session_on_enter(session=session3, user=user)
+ event_manager.session_on_enter(session=session3, user=user2)
session3.update(
finished=utc_now,
status=Status.FAIL,
status_code_1=StatusCode1.BUYER_FAIL,
)
- event_manager.session_on_finish(session=session3, user=user)
+ event_manager.session_on_finish(session=session3, user=user2)
avg_loi_complete = (
round(session.elapsed.total_seconds())
+ round(session2.elapsed.total_seconds())
) / 2
- assert event_manager.get_session_stats(product_id) == {
+ assert event_manager.get_global_session_stats() == {
"session_enters_last_1h": 2,
"session_enters_last_24h": 3,
"session_fails_last_1h": 1,
@@ -307,7 +320,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),
@@ -321,7 +334,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,
@@ -384,7 +397,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 +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()
@@ -414,7 +427,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 +441,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()
@@ -444,14 +457,15 @@ class TestTaskStatsManager:
class TestChannelsSubscriptions:
+ @pytest.mark.skip("sits there doing nothing forever? todo")
def test_stats_worker(
self,
- event_manager,
- event_subscriber,
- product_id,
- user_factory,
- utc_hour_ago,
- utc_now,
+ event_manager: EventManager,
+ event_subscriber: EventSubscriber,
+ product_id: str,
+ user_factory: Callable[..., User],
+ utc_hour_ago: datetime,
+ utc_now: datetime,
):
event_manager.clear_stats()
assert event_subscriber.pubsub
@@ -481,7 +495,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,
@@ -496,8 +510,8 @@ class TestChannelsSubscriptions:
wall.update(
status=Status.COMPLETE,
status_code_1=StatusCode1.COMPLETE,
- finished=datetime.now(tz=timezone.utc),
- cpi=Decimal("1"),
+ finished=datetime.now(tz=UTC),
+ cpi=Decimal(1),
)
event_manager.handle_task_finish(wall, session, user)
msg = event_subscriber.get_next_message()
diff --git a/tests/managers/test_lucid.py b/tests/managers/test_lucid.py
index 1a1bae7..6771a0c 100644
--- a/tests/managers/test_lucid.py
+++ b/tests/managers/test_lucid.py
@@ -1,14 +1,21 @@
+from __future__ import annotations
+
+from typing import TYPE_CHECKING
+
import pytest
from generalresearch.managers.lucid.profiling import get_profiling_library
+if TYPE_CHECKING:
+ from generalresearch.pg_helper import PostgresConfig
+
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, pks=pks)
assert len(qids) == len(qs)
diff --git a/tests/managers/test_userpid.py b/tests/managers/test_userpid.py
index 4a3f699..e74e40b 100644
--- a/tests/managers/test_userpid.py
+++ b/tests/managers/test_userpid.py
@@ -1,11 +1,12 @@
+from __future__ import annotations
+
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 +14,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_buyer.py b/tests/managers/thl/test_buyer.py
index 69ea105..0ab2d52 100644
--- a/tests/managers/thl/test_buyer.py
+++ b/tests/managers/thl/test_buyer.py
@@ -1,14 +1,24 @@
-from generalresearch.models import Source
+from __future__ import annotations
+
+from collections.abc import Callable
+from typing import TYPE_CHECKING
+
+from generalresearch.models.definitions import Source
+
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.buyer import BuyerManager
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..fc364f2 100644
--- a/tests/managers/thl/test_cashout_method.py
+++ b/tests/managers/thl/test_cashout_method.py
@@ -1,62 +1,75 @@
+from __future__ import annotations
+
+from collections.abc import Callable
+from typing import TYPE_CHECKING
+
import pytest
-from generalresearch.models.thl.wallet import PayoutType
from generalresearch.models.thl.wallet.cashout_method import (
CashMailCashoutMethodData,
PaypalCashoutMethodData,
USDeliveryAddress,
)
-from test_utils.managers.cashout_methods import (
- EXAMPLE_TANGO_CASHOUT_METHODS,
-)
+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:
- 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],
+ example_tango_cashout_methods: list[CashoutMethod],
+ ):
+ 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]
- assert EXAMPLE_TANGO_CASHOUT_METHODS[0] == cm
+ 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
-class TestAMTCashoutMethods:
-
- def test_create_and_get(self, cashout_method_manager, 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 = [x for x in res if x.name == "AMT Bonus"][0]
- assert AMT_BONUS_CASHOUT_METHOD == cm
-
- def test_user(
- self, cashout_method_manager, user_with_wallet_amt, 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
-
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 +108,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..a2805bc 100644
--- a/tests/managers/thl/test_category.py
+++ b/tests/managers/thl/test_category.py
@@ -1,12 +1,21 @@
+from __future__ import annotations
+
+from collections.abc import Callable
+from typing import TYPE_CHECKING
+
import pytest
from generalresearch.models.thl.category import Category
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.category import CategoryManager
+ from generalresearch.pg_helper import PostgresConfig
+
class TestCategory:
@pytest.fixture
- def beauty_fitness(self, thl_web_rw):
+ def beauty_fitness(self) -> Category:
return Category(
uuid="12c1e96be82c4642a07a12a90ce6f59e",
@@ -16,72 +25,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.path}/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_contest/test_leaderboard.py b/tests/managers/thl/test_contest/test_leaderboard.py
index 80a88a5..8aa0780 100644
--- a/tests/managers/thl/test_contest/test_leaderboard.py
+++ b/tests/managers/thl/test_contest/test_leaderboard.py
@@ -1,34 +1,41 @@
-from datetime import datetime, timezone, timedelta
+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.models.thl.contest.definitions import (
- ContestStatus,
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 test_utils.managers.contest.conftest import (
- leaderboard_contest_in_db as contest_in_db,
- leaderboard_contest_create as contest_create,
)
+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:
+class TestLeaderboardContestCRUD:
def test_create(
self,
- contest_create: LeaderboardContestCreate,
+ leaderboard_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
+ product_id=product_user_wallet_yes.uuid,
+ contest_create=leaderboard_contest_create,
)
c_out = contest_manager.get(c.uuid)
assert c == c_out
@@ -39,18 +46,19 @@ 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,
user_with_wallet: User,
- contest_in_db: LeaderboardContest,
- thl_lm,
- contest_manager,
- user_manager,
- thl_redis,
+ leaderboard_contest_in_db: LeaderboardContest,
+ thl_ledger_manager: ThlLedgerManager,
+ contest_manager: ContestManager,
+ user_manager: UserManager,
+ thl_redis_config: RedisConfig,
):
- contest = contest_in_db
+ thl_redis = thl_redis_config.create_redis_client()
+ contest = leaderboard_contest_in_db
user = user_with_wallet
c: LeaderboardContest = contest_manager.get(contest_uuid=contest.uuid)
@@ -77,14 +85,15 @@ class TestLeaderboardContestCRUD:
def test_contest_ends(
self,
user_with_wallet: User,
- contest_in_db: LeaderboardContest,
- thl_lm,
- contest_manager,
- user_manager,
- thl_redis,
+ leaderboard_contest_in_db: LeaderboardContest,
+ thl_ledger_manager: ThlLedgerManager,
+ contest_manager: ContestManager,
+ user_manager: UserManager,
+ thl_redis_config: RedisConfig,
):
+ thl_redis = thl_redis_config.create_redis_client()
# The contest should be over. We need to trigger it.
- contest = contest_in_db
+ contest = leaderboard_contest_in_db
contest._redis_client = thl_redis
contest._user_manager = user_manager
user = user_with_wallet
@@ -100,18 +109,22 @@ 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()
assert decision
assert reason == ContestEndReason.ENDS_AT
- contest_manager.end_contest_if_over(contest=contest, ledger_manager=thl_lm)
+ contest_manager.end_contest_if_over(
+ contest=contest, ledger_manager=thl_ledger_manager
+ )
c: LeaderboardContest = contest_manager.get(contest_uuid=contest.uuid)
assert c.status == ContestStatus.COMPLETED
@@ -129,10 +142,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 7312a64..f26819b 100644
--- a/tests/managers/thl/test_contest/test_milestone.py
+++ b/tests/managers/thl/test_contest/test_milestone.py
@@ -1,34 +1,41 @@
-from datetime import datetime, timezone
+from __future__ import annotations
+
+from collections.abc import Callable
+from datetime import UTC, datetime
+from typing import TYPE_CHECKING
from generalresearch.models.thl.contest.definitions import (
- ContestStatus,
ContestEndReason,
+ ContestEntryTrigger,
+ ContestStatus,
)
from generalresearch.models.thl.contest.milestone import (
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,
- milestone_contest_create as contest_create,
- milestone_contest_factory as contest_factory,
)
+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.product import Product
+ from generalresearch.models.thl.user import User
-class TestMilestoneContest:
- def test_should_end(self, contest: MilestoneContest, thl_lm, contest_manager):
+class TestMilestoneContest:
+ def test_should_end(
+ self,
+ milestone_contest: MilestoneContest,
+ ):
+ contest = milestone_contest
# contest is active and has no entries
should, msg = contest.should_end()
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
@@ -43,16 +50,15 @@ class TestMilestoneContest:
class TestMilestoneContestCRUD:
-
def test_create(
self,
- contest_create: MilestoneContestCreate,
+ milestone_contest_create: MilestoneContestCreate,
product_user_wallet_yes: Product,
- thl_lm,
- contest_manager,
+ contest_manager: ContestManager,
):
c = contest_manager.create(
- product_id=product_user_wallet_yes.uuid, contest_create=contest_create
+ product_id=product_user_wallet_yes.uuid,
+ contest_create=milestone_contest_create,
)
c_out = contest_manager.get(c.uuid)
assert c == c_out
@@ -68,20 +74,20 @@ class TestMilestoneContestCRUD:
def test_enter(
self,
user_with_wallet: User,
- contest_in_db: MilestoneContest,
- thl_lm,
- contest_manager,
+ milestone_contest_in_db: MilestoneContest,
+ 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.
- contest = contest_in_db
+ contest = milestone_contest_in_db
user = user_with_wallet
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,
)
@@ -96,17 +102,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(
@@ -122,21 +130,21 @@ class TestMilestoneContestCRUD:
def test_enter_win(
self,
user_with_wallet: User,
- contest_in_db: MilestoneContest,
- thl_lm,
- contest_manager,
+ milestone_contest_in_db: MilestoneContest,
+ 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
+ contest = milestone_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
@@ -151,7 +159,7 @@ class TestMilestoneContestCRUD:
contest_uuid=contest.uuid,
user=user,
country_iso="us",
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
incr=1,
)
@@ -171,9 +179,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
@@ -182,22 +193,22 @@ 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,
+ milestone_contest_in_db: MilestoneContest,
+ 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)]
- contest = contest_in_db
+ contest = milestone_contest_in_db
for u in users:
contest_manager.enter_milestone_contest(
contest_uuid=contest.uuid,
user=u,
country_iso="us",
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
incr=3,
)
@@ -208,29 +219,33 @@ class TestMilestoneContestCRUD:
def test_trigger(
self,
user_with_wallet: User,
- contest_in_db: MilestoneContest,
- thl_lm,
- contest_manager,
+ milestone_contest_in_db: MilestoneContest,
+ 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
# Assert this contest got entered
c: MilestoneUserView = contest_manager.get_milestone_user_view(
- contest_uuid=contest_in_db.uuid, user=user_with_wallet
+ contest_uuid=milestone_contest_in_db.uuid, user=user_with_wallet
)
assert c.user_amount == 1
class TestMilestoneContestUserViews:
def test_list_user_eligible_country(
- self, user_with_wallet: User, contest_factory, thl_lm, contest_manager
+ self,
+ user_with_wallet: User,
+ milestone_contest_factory: Callable[..., MilestoneContest],
+ thl_ledger_manager: ThlLedgerManager,
+ contest_manager: ContestManager,
):
# No contests exists
cs = contest_manager.get_many_by_user_eligible(
@@ -239,7 +254,7 @@ class TestMilestoneContestUserViews:
assert len(cs) == 0
# Create a contest. It'll be in the US/CA
- contest_factory(country_isos={"us", "ca"})
+ milestone_contest_factory(country_isos={"us", "ca"})
# Not eligible in mexico
cs = contest_manager.get_many_by_user_eligible(
@@ -252,7 +267,7 @@ class TestMilestoneContestUserViews:
assert len(cs) == 1
# Create another, any country
- contest_factory(country_isos=None)
+ milestone_contest_factory(country_isos=None)
cs = contest_manager.get_many_by_user_eligible(
user=user_with_wallet, country_iso="mx"
)
@@ -263,10 +278,14 @@ 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,
+ milestone_contest_factory: Callable[..., MilestoneContest],
+ thl_ledger_manager: ThlLedgerManager,
+ contest_manager: ContestManager,
):
# User reaches milestone after 1 complete
- c = contest_factory(target_amount=1)
+ c = milestone_contest_factory(target_amount=1)
user = user_with_money
cs = contest_manager.get_many_by_user_eligible(
@@ -275,7 +294,10 @@ class TestMilestoneContestUserViews:
assert len(cs) == 1
contest_manager.enter_milestone_contest(
- contest_uuid=c.uuid, user=user, country_iso="us", ledger_manager=thl_lm
+ contest_uuid=c.uuid,
+ user=user,
+ country_iso="us",
+ ledger_manager=thl_ledger_manager,
)
# User isn't eligible anymore
diff --git a/tests/managers/thl/test_contest/test_raffle.py b/tests/managers/thl/test_contest/test_raffle.py
index 060055a..7388991 100644
--- a/tests/managers/thl/test_contest/test_raffle.py
+++ b/tests/managers/thl/test_contest/test_raffle.py
@@ -1,4 +1,8 @@
-from datetime import datetime, timezone
+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
@@ -9,44 +13,49 @@ 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,
-)
-from generalresearch.models.thl.contest.exceptions import ContestError
-from generalresearch.models.thl.contest.raffle import (
+from generalresearch.models.thl.contest.contest_entry import (
ContestEntry,
ContestEntryType,
)
-from generalresearch.models.thl.contest.raffle import (
- RaffleContest,
- RaffleContestCreate,
- RaffleUserView,
-)
-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,
- raffle_contest_create as contest_create,
- raffle_contest_factory as contest_factory,
+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 RaffleContest
+
+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:
- def test_should_end(self, contest: RaffleContest, thl_lm, contest_manager):
+ def test_should_end(
+ self,
+ raffle_contest: RaffleContest,
+ ):
+ contest = raffle_contest
# contest is active and has no entries
should, msg = contest.should_end()
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
@@ -63,13 +72,12 @@ class TestRaffleContestCRUD:
def test_create(
self,
- contest_create: RaffleContestCreate,
+ raffle_contest_create: RaffleContestCreate,
product_user_wallet_yes: Product,
- thl_lm,
- contest_manager,
+ contest_manager: ContestManager,
):
c = contest_manager.create(
- product_id=product_user_wallet_yes.uuid, contest_create=contest_create
+ product_id=product_user_wallet_yes.uuid, contest_create=raffle_contest_create
)
c_out = contest_manager.get(c.uuid)
assert c == c_out
@@ -85,18 +93,20 @@ class TestRaffleContestCRUD:
def test_enter(
self,
user_with_money: User,
- contest_in_db: RaffleContest,
- thl_lm,
- contest_manager,
+ raffle_contest_in_db: RaffleContest,
+ thl_ledger_manager: ThlLedgerManager,
+ contest_manager: ContestManager,
):
# Raffle ends at $1.00. User enters for $0.60
print(user_with_money.product_id)
- print(contest_in_db.product_id)
- print(contest_in_db.uuid)
- contest = contest_in_db
+ print(raffle_contest_in_db.product_id)
+ print(raffle_contest_in_db.uuid)
+ contest = raffle_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)
@@ -105,7 +115,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)
@@ -120,30 +130,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,
+ raffle_contest_in_db: RaffleContest,
+ 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
+ contest = raffle_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(
@@ -155,7 +170,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
@@ -175,25 +190,33 @@ 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,
+ raffle_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
- c = contest_factory(
+ c = raffle_contest_factory(
prizes=[
ContestPrize(
name="$1.00 bonus",
@@ -205,12 +228,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(
@@ -220,28 +245,35 @@ 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,
+ raffle_contest_in_db: RaffleContest,
+ thl_ledger_manager: ThlLedgerManager,
+ contest_manager: ContestManager,
):
- c = contest_in_db
+ c = raffle_contest_in_db
user = user_with_wallet
# Tries to enter $0
@@ -260,7 +292,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"
@@ -271,16 +303,20 @@ 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,
+ raffle_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(
+ c = raffle_contest_factory(
entry_rule=ContestEntryRule(
max_entry_amount_per_user=USDCent(10),
max_daily_entries_per_user=USDCent(8),
@@ -299,7 +335,7 @@ class TestRaffleContestCRUD:
contest_uuid=c.uuid,
entry=entry,
country_iso="us",
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
)
assert "Entry would exceed max amount per user." in str(e.value)
@@ -312,7 +348,7 @@ class TestRaffleContestCRUD:
contest_uuid=c.uuid,
entry=entry,
country_iso="us",
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
)
assert "Entry would exceed max amount per user per day." in str(e.value)
@@ -324,7 +360,7 @@ class TestRaffleContestCRUD:
contest_uuid=c.uuid,
entry=entry,
country_iso="us",
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
)
# Then can't anymore
@@ -336,14 +372,18 @@ class TestRaffleContestCRUD:
contest_uuid=c.uuid,
entry=entry,
country_iso="us",
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
)
assert "Entry would exceed max amount per user per day." in str(e.value)
class TestRaffleContestUserViews:
def test_list_user_eligible_country(
- self, user_with_wallet: User, contest_factory, thl_lm, contest_manager
+ self,
+ user_with_wallet: User,
+ raffle_contest_factory: Callable[..., Contest],
+ thl_ledger_manager: ThlLedgerManager,
+ contest_manager: ContestManager,
):
# No contests exists
cs = contest_manager.get_many_by_user_eligible(
@@ -352,7 +392,7 @@ class TestRaffleContestUserViews:
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(
@@ -365,7 +405,7 @@ class TestRaffleContestUserViews:
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"
)
@@ -376,9 +416,13 @@ 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,
+ raffle_contest_factory: Callable[..., Contest],
+ thl_ledger_manager: ThlLedgerManager,
+ contest_manager: ContestManager,
):
- c = contest_factory(
+ c = raffle_contest_factory(
end_condition=ContestEndCondition(target_entry_amount=USDCent(10)),
entry_rule=ContestEntryRule(
max_entry_amount_per_user=USDCent(1),
@@ -398,7 +442,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
@@ -422,9 +466,13 @@ 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,
+ raffle_contest_factory: Callable[..., Contest],
+ thl_ledger_manager: ThlLedgerManager,
+ contest_manager: ContestManager,
):
- c = contest_factory(
+ c = raffle_contest_factory(
end_condition=ContestEndCondition(target_entry_amount=USDCent(100)),
)
entry = ContestEntry(
@@ -436,7 +484,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)
@@ -458,9 +506,13 @@ 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,
+ raffle_contest_factory: Callable[..., Contest],
+ thl_ledger_manager: ThlLedgerManager,
+ contest_manager: ContestManager,
):
- c = contest_factory(entry_type=ContestEntryType.COUNT)
+ c = raffle_contest_factory(entry_type=ContestEntryType.COUNT)
entry = ContestEntry(
entry_type=ContestEntryType.COUNT,
user=user_with_wallet,
@@ -470,5 +522,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_harmonized_uqa.py b/tests/managers/thl/test_harmonized_uqa.py
index 6bbbbe1..2fc0ff0 100644
--- a/tests/managers/thl/test_harmonized_uqa.py
+++ b/tests/managers/thl/test_harmonized_uqa.py
@@ -1,13 +1,18 @@
-from datetime import datetime, timezone
+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 (
- UserQuestionAnswer,
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")
@@ -18,7 +23,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 +43,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 +62,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 +108,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_ipinfo.py b/tests/managers/thl/test_ipinfo.py
index 847b00c..c021eb9 100644
--- a/tests/managers/thl/test_ipinfo.py
+++ b/tests/managers/thl/test_ipinfo.py
@@ -1,51 +1,75 @@
+from collections.abc import Callable
+from typing import TYPE_CHECKING
+
import faker
from generalresearch.managers.thl.ipinfo import (
+ GeoIpInfoManager,
IPGeonameManager,
IPInformationManager,
- GeoIpInfoManager,
)
-from generalresearch.models.thl.ipinfo import IPGeoname, IPInformation
+from generalresearch.models.thl.ipinfo import (
+ GeoIPInformation,
+ IPGeoname,
+ IPInformation,
+)
+
+if TYPE_CHECKING:
+ from generalresearch.pg_helper import PostgresConfig
+ from generalresearch.redis_helper import RedisConfig
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)
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()
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)
- def test_create(self, ip_information_manager: IPInformationManager):
- instance = ip_information_manager.create_dummy()
-
+ def test_create(
+ self,
+ ip_information_factory: Callable[..., IPInformation],
+ ip_information_manager: IPInformationManager,
+ ):
+ instance = ip_information_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(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
@@ -57,13 +81,21 @@ 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)
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]
@@ -90,7 +122,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"
@@ -105,13 +142,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_accounts.py b/tests/managers/thl/test_ledger/test_lm_accounts.py
index 5cfaac1..3af10e7 100644
--- a/tests/managers/thl/test_ledger/test_lm_accounts.py
+++ b/tests/managers/thl/test_ledger/test_lm_accounts.py
@@ -1,9 +1,12 @@
+from __future__ import annotations
+
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
+from pydantic import PositiveInt
from generalresearch.currency import LedgerCurrency
from generalresearch.managers.base import Permission
@@ -11,6 +14,7 @@ from generalresearch.managers.thl.ledger_manager.exceptions import (
LedgerAccountDoesntExistError,
)
from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager
+from generalresearch.models.custom_types import UUIDStr
from generalresearch.models.thl.ledger import (
AccountType,
Direction,
@@ -19,22 +23,10 @@ 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
- from generalresearch.models.thl import Direction
from generalresearch.models.thl.ledger import (
- AccountType,
- 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(
@@ -51,53 +43,63 @@ class TestLedgerAccountManagerNoResults:
def test_get_account_no_results(
self,
- currency: "LedgerCurrency",
+ currency: LedgerCurrency,
kind: str,
- acct_id: "UUIDStr",
- lm: "LedgerManager",
+ acct_id: UUIDStr,
+ 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
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
- 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",
+ currency: LedgerCurrency,
kind: str,
- acct_id: "UUIDStr",
- lm: "LedgerManager",
+ acct_id: UUIDStr,
+ ledger_manager: LedgerManager,
):
- qn = ":".join([currency, kind, acct_id])
+ qn = f"{currency}:{kind}:{acct_id}"
# (1) .get_many_
- assert lm.get_account_many_(qualified_names=[qn], raise_on_error=False) == []
+ 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(
@@ -114,10 +116,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,
+ ledger_manager: LedgerManager,
):
"""Confirm that the Permission values that are set on the Ledger Manger
allow the Creation action to occur.
@@ -134,11 +136,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:
@@ -149,11 +151,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:
@@ -164,10 +166,10 @@ class TestLedgerAccountManagerCreate:
def test_create(
self,
- currency: "LedgerCurrency",
- account_type: "AccountType",
- direction: "Direction",
- lm: "LedgerManager",
+ currency: LedgerCurrency,
+ account_type: AccountType,
+ direction: Direction,
+ ledger_manager: LedgerManager,
):
"""Confirm that the Permission values that are set on the Ledger Manger
allow the Creation action to occur.
@@ -184,20 +186,20 @@ 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
def test_get_or_create(
self,
- currency: "LedgerCurrency",
- account_type: "AccountType",
- direction: "Direction",
- lm: "LedgerManager",
+ currency: LedgerCurrency,
+ account_type: AccountType,
+ direction: Direction,
+ ledger_manager: LedgerManager,
):
"""Confirm that the Permission values that are set on the Ledger Manger
allow the Creation action to occur.
@@ -214,27 +216,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
@@ -243,30 +249,30 @@ 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,
+ 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))
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,
+ ledger_manager: LedgerManager,
):
"""Try searching for random metadata and confirm it's always 0 because
Tx can be found.
@@ -275,7 +281,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
@@ -285,7 +291,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,
@@ -302,7 +308,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,
@@ -311,7 +317,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,
@@ -320,7 +326,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_lm_tx.py b/tests/managers/thl/test_ledger/test_lm_tx.py
index 37b7ba3..025f6ac 100644
--- a/tests/managers/thl/test_ledger/test_lm_tx.py
+++ b/tests/managers/thl/test_ledger/test_lm_tx.py
@@ -1,33 +1,41 @@
+from __future__ import annotations
+
from decimal import Decimal
from random import randint
+from typing import TYPE_CHECKING
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,
LedgerEntry,
LedgerTransaction,
)
+if TYPE_CHECKING:
+ from generalresearch.models.thl.ledger import (
+ LedgerAccount,
+ )
+
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.
"""
- acct_uuid = uuid4().hex
-
# (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 +45,12 @@ class TestLedgerManagerCreateTx:
== "LedgerTransactionManager has insufficient Permissions"
)
- def test_create_assertions(self, ledger_account_debit, ledger_account_credit, lm):
+ def test_create_assertions(
+ self,
+ ledger_manager: LedgerManager,
+ ):
with pytest.raises(expected_exception=ValueError) as excinfo:
- lm.create_tx(
+ ledger_manager.create_tx(
entries=[
{
"direction": Direction.CREDIT,
@@ -53,7 +64,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 +86,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 +115,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(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 +137,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 +158,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 +219,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..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,25 +1,41 @@
-from generalresearch.models.thl.ledger import LedgerEntry
+from __future__ import annotations
+
+from typing import TYPE_CHECKING
+
+from generalresearch.models.thl.ledger import (
+ LedgerEntry,
+)
+
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager
+ from generalresearch.models.thl.ledger import (
+ 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 df2611b..166598e 100644
--- a/tests/managers/thl/test_ledger/test_lm_tx_locks.py
+++ b/tests/managers/thl/test_ledger/test_lm_tx_locks.py
@@ -1,29 +1,38 @@
+from __future__ import annotations
+
import logging
-from datetime import datetime, timezone, timedelta
+from collections.abc import Callable, Generator
+from datetime import UTC, datetime, timedelta
from decimal import Decimal
-from typing import Callable
+from typing import TYPE_CHECKING
import pytest
+from pytest import LogCaptureFixture
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.definitions 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
+
+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")
@@ -32,17 +41,17 @@ class TestLedgerLocks:
def test_a(
self,
- user_factory,
- session_factory,
- product_user_wallet_no,
- create_main_accounts,
- caplog,
- thl_lm,
- lm,
- utc_hour_ago,
- currency,
- wall_factory,
- delete_ledger_db,
+ user_factory: Callable[..., User],
+ session_factory: Callable[..., Session],
+ product_user_wallet_no: Product,
+ create_main_accounts: Callable[..., None],
+ caplog: Generator[LogCaptureFixture],
+ 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.
@@ -56,18 +65,22 @@ 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],
)
# 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
@@ -77,7 +90,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
@@ -86,60 +99,63 @@ 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")
- with caplog.at_level(logging.ERROR):
- with pytest.raises(expected_exception=LedgerTransactionCreateLockError):
- tx = thl_lm.create_tx_protected(
- lock_key=lock_key,
- condition=condition,
- create_tx_func=create_tx_func,
- )
- assert tx is None
+ ledger_manager.redis_client.set(lock_name, "1")
+ 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
- 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,
- caplog,
- thl_lm,
- lm,
+ user_factory: Callable[..., User],
+ product_user_wallet_no: Product,
+ create_main_accounts: Callable[..., None],
+ delete_ledger_db: Callable[..., None],
+ caplog: Generator[LogCaptureFixture],
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
):
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
@@ -155,7 +171,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
@@ -170,7 +188,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
@@ -179,7 +199,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
@@ -187,7 +207,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,
@@ -196,52 +216,53 @@ 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")
- with caplog.at_level(logging.DEBUG):
- with pytest.raises(expected_exception=LedgerTransactionCreateLockError):
- tx = thl_lm.create_tx_task_complete(
- wall=wall3, user=user, created=wall3.started
- )
- assert isinstance(tx, LedgerTransaction)
+ ledger_manager.redis_client.set(name=lock_name, value="1")
+ 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
- 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
@@ -250,29 +271,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(
@@ -283,39 +309,46 @@ 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)
+ 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)
@@ -327,45 +360,49 @@ 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)
+ 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..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,34 +1,55 @@
+from __future__ import annotations
+
+from typing import TYPE_CHECKING
+
+if TYPE_CHECKING:
+ 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 01d5fe1..107ff00 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,41 @@
+from __future__ import annotations
+
+from collections.abc import Callable
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
+from generalresearch.currency import LedgerCurrency
+from generalresearch.managers.thl.ledger_manager.exceptions import (
+ LedgerAccountDoesntExistError,
+)
+from generalresearch.models.thl.ledger import (
+ AccountType,
+ Direction,
+ LedgerAccount,
+)
+from generalresearch.models.thl.product import Product
+
+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:
- 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,
- )
+ 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 +47,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 (
- LedgerAccount,
- Direction,
- AccountType,
- )
+ 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 +72,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 (
- Direction,
- AccountType,
- )
+ 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 +98,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 (
- Direction,
- AccountType,
- )
-
- 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 +124,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 (
- Direction,
- AccountType,
- )
+ 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 +150,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 +169,79 @@ 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, create_main_accounts
):
from generalresearch.models.thl.ledger import (
- LedgerAccount,
AccountType,
+ LedgerAccount,
)
- res = thl_lm.get_account_task_complete_revenue()
+ create_main_accounts()
+
+ 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,
+ create_main_accounts
+ ):
+ create_main_accounts()
from generalresearch.models.thl.ledger import (
- LedgerAccount,
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.models.thl.user import User
- from generalresearch.managers.thl.ledger_manager.exceptions import (
- LedgerAccountDoesntExistError,
- )
+ 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 isinstance(res, LedgerAccount)
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 +249,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 +286,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 +298,50 @@ class TestThlLedgerManagerAccounts:
assert len(res) == 2
# Confirm an empty array comes back for all unknown qualified names
- res = lm.get_accounts_if_exists(
+ assert isinstance(ledger_manager.currency, LedgerCurrency)
+ 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])
+ 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 (
- LedgerAccount,
- Direction,
- AccountType,
- )
+ def test_get_or_create(self, ledger_account_manager: LedgerAccountManager):
u = uuid4().hex
name = f"test-{u[:8]}"
@@ -297,48 +359,52 @@ class TestLedgerAccountManager:
# First we want to validate that using the get_account method raises
# an error for a random LedgerAccount which we know does not exist.
with pytest.raises(LedgerAccountDoesntExistError):
- lam.get_account(qualified_name=account.qualified_name)
+ ledger_account_manager.get_account(qualified_name=account.qualified_name)
# Now that we know it doesn't exist, get_or_create for it
- instance = lam.get_account_or_create(account=account)
+ instance = ledger_account_manager.get_account_or_create(account=account)
# It should always return
assert isinstance(instance, LedgerAccount)
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 (
- LedgerAccount,
- AccountType,
- )
+ def test_get(
+ self,
+ user: User,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_account_manager: LedgerAccountManager,
+ ):
+
+ assert isinstance(user.product, Product)
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.models.thl.product import Product
- from generalresearch.managers.thl.ledger_manager.exceptions import (
- LedgerAccountDoesntExistError,
- )
-
+ 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 +412,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 +422,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 +435,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 +445,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 (
- LedgerAccount,
- Direction,
- AccountType,
- )
+ def test_create_account(self, ledger_account_manager: LedgerAccountManager):
u = uuid4().hex
name = f"test-{u[:8]}"
@@ -406,6 +467,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 294d092..27ddc29 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,7 +1,11 @@
+from __future__ import annotations
+
import logging
-from datetime import datetime, timezone, timedelta
+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
@@ -11,27 +15,35 @@ 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.models import Source
+from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager
+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.session import (
- Wall,
+ Session,
Status,
StatusCode1,
- Session,
+ Wall,
)
-from generalresearch.models.thl.user import User
-from generalresearch.models.thl.wallet import PayoutType
from generalresearch.redis_helper import RedisConfig
+if TYPE_CHECKING:
+ from generalresearch.currency import LedgerCurrency
+ from generalresearch.managers.thl.payout import (
+ BrokerageProductPayoutEventManager,
+ BusinessPayoutEventManager,
+ )
+ 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")
@@ -42,20 +54,18 @@ def broken_release(self, *args, **kwargs):
class TestThlLedgerManagerBPPayout:
+ @pytest.fixture(autouse=True)
+ def setup(self, create_main_accounts):
+ create_main_accounts()
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,
caplog,
- thl_lm,
- delete_ledger_db,
+ thl_ledger_manager: ThlLedgerManager,
):
- 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(
@@ -69,31 +79,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,
@@ -101,7 +109,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),
@@ -109,13 +117,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),
@@ -125,19 +135,21 @@ class TestThlLedgerManagerBPPayout:
)
payoutevent_uuid = uuid4().hex
- with caplog.at_level(logging.INFO):
- with pytest.raises(LedgerTransactionConditionFailedError):
- thl_lm.create_tx_bp_payout(
- user.product,
- amount=USDCent(10_000),
- created=now + timedelta(minutes=2),
- skip_one_per_day_check=True,
- skip_wallet_balance_check=False,
- payoutevent_uuid=payoutevent_uuid,
- )
+ with (
+ caplog.at_level(logging.INFO),
+ pytest.raises(LedgerTransactionConditionFailedError),
+ ):
+ thl_ledger_manager.create_tx_bp_payout(
+ user.product,
+ amount=USDCent(10_000),
+ created=now + timedelta(minutes=2),
+ skip_one_per_day_check=True,
+ skip_wallet_balance_check=False,
+ payoutevent_uuid=payoutevent_uuid,
+ )
assert "failed condition check balance:" in caplog.text
- thl_lm.create_tx_bp_payout(
+ thl_ledger_manager.create_tx_bp_payout(
product=user.product,
amount=USDCent(10_00),
created=now + timedelta(minutes=2),
@@ -145,20 +157,26 @@ 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,
- 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,
@@ -177,31 +195,38 @@ 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(
- product=product,
- amount=rand_amount,
- payoutevent_uuid=uuid4().hex,
- created=datetime.now(tz=timezone.utc),
- skip_wallet_balance_check=False,
- skip_one_per_day_check=False,
- skip_flag_check=False,
- )
+ with (
+ caplog.at_level(logging.INFO),
+ pytest.raises(expected_exception=LedgerTransactionConditionFailedError),
+ ):
+ thl_ledger_manager.create_tx_bp_payout(
+ product=product,
+ amount=rand_amount,
+ payoutevent_uuid=uuid4().hex,
+ created=datetime.now(tz=UTC),
+ skip_wallet_balance_check=False,
+ skip_one_per_day_check=False,
+ skip_flag_check=False,
+ )
assert "failed condition check >1 tx per day" in caplog.text
- def test_create_tx_redis_failure(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=timezone.utc)
+ 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
)
@@ -222,43 +247,49 @@ 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,
- created=datetime.now(tz=timezone.utc),
+ created=datetime.now(tz=UTC),
)
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=timezone.utc)
+ 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,
- created=datetime.now(tz=timezone.utc),
+ created=datetime.now(tz=UTC),
)
# 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,
- created=datetime.now(tz=timezone.utc),
+ created=datetime.now(tz=UTC),
)
assert e.type is LedgerTransactionFlagAlreadyExistsError
@@ -266,251 +297,261 @@ 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(
+ thl_ledger_manager.create_tx_bp_payout(
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"
# 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,
- 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):
+ def test_create_tx_redis_lock_release_error(
+ self,
+ product: Product,
+ thl_ledger_manager: ThlLedgerManager,
+ monkeypatch: pytest.MonkeyPatch,
+ ):
rand_amount: USDCent = USDCent(randint(100, 1_000))
payoutevent_uuid = uuid4().hex
- now = datetime.now(tz=timezone.utc)
- bp_wallet_account = thl_lm.get_account_or_create_bp_wallet(product=product)
+ now = datetime.now(tz=UTC)
+ 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
)
- original_acquire = Lock.acquire
- original_release = Lock.release
- Lock.acquire = broken_acquire
-
# 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(
- product=product,
- amount=rand_amount,
- payoutevent_uuid=payoutevent_uuid,
- created=datetime.now(tz=timezone.utc),
- )
+ with monkeypatch.context() as m:
+ m.setattr(Lock, "acquire", broken_acquire)
+ with pytest.raises(expected_exception=Exception) as e:
+ thl_ledger_manager.create_tx_bp_payout(
+ product=product,
+ amount=rand_amount,
+ payoutevent_uuid=payoutevent_uuid,
+ created=datetime.now(tz=UTC),
+ )
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
- Lock.acquire = original_acquire
- Lock.release = broken_release
-
# 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(
- product=product,
- amount=rand_amount,
- payoutevent_uuid=payoutevent_uuid,
- created=datetime.now(tz=timezone.utc),
- )
- assert e.type is LedgerTransactionReleaseLockError
+ with monkeypatch.context() as m:
+ m.setattr(Lock, "release", broken_release)
+ with pytest.raises(LedgerTransactionReleaseLockError) as e:
+ thl_ledger_manager.create_tx_bp_payout(
+ product=product,
+ amount=rand_amount,
+ payoutevent_uuid=payoutevent_uuid,
+ created=datetime.now(tz=UTC),
+ )
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
class TestPayoutEventManagerBPPayout:
+ @pytest.fixture(autouse=True)
+ def setup(self, create_main_accounts):
+ create_main_accounts()
- 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,
+ business_payout_event_manager: BusinessPayoutEventManager,
+ ):
rand_amount: USDCent = USDCent(randint(100, 1_000))
- now = datetime.now(tz=timezone.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(
+ now = datetime.now(tz=UTC)
+ 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
- brokerage_product_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
+ assert thl_ledger_manager.get_account_balance(bp_wallet_account) == rand_amount
- pe = brokerage_product_payout_event_manager.create_bp_payout_event(
- thl_ledger_manager=thl_lm,
+ bpe = business_payout_event_manager.create_bp_payout_event(
+ thl_ledger_manager=thl_ledger_manager,
product=product,
created=now,
amount=rand_amount,
- payout_type=PayoutType.ACH,
+ ext_ref_id=uuid4().hex,
)
+ bp_pe = bpe.bp_payouts[0]
assert brokerage_product_payout_event_manager.check_for_ledger_tx(
- thl_ledger_manager=thl_lm,
- product_id=product.id,
- amount=rand_amount,
- payout_event=pe,
+ thl_ledger_manager=thl_ledger_manager,
+ payout_event=bp_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,
+ business_payout_event_manager: BusinessPayoutEventManager,
+ monkeypatch: pytest.MonkeyPatch,
):
caplog.set_level("WARNING")
- original_acquire = Lock.acquire
- original_release = Lock.release
+ ext_ref_id = uuid4().hex
rand_amount: USDCent = USDCent(randint(100, 1_000))
- now = datetime.now(tz=timezone.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
+ now = datetime.now(tz=UTC)
+ 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, created=now, direction=Direction.CREDIT
)
- assert thl_lm.get_account_balance(bp_wallet_account) == rand_amount
- brokerage_product_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
+ assert thl_ledger_manager.get_account_balance(bp_wallet_account) == rand_amount
# 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,
- product=product,
- created=now,
- amount=rand_amount,
- payout_type=PayoutType.ACH,
- )
- assert e.type is LedgerTransactionCreateError
+ with monkeypatch.context() as m:
+ m.setattr(Lock, "acquire", broken_acquire)
+ with pytest.raises(LedgerTransactionCreateError) as e:
+ business_payout_event_manager.create_bp_payout_event(
+ thl_ledger_manager=thl_ledger_manager,
+ product=product,
+ created=now,
+ amount=rand_amount,
+ ext_ref_id=ext_ref_id,
+ )
assert str(e.value) == "Redis error: Simulated timeout during acquire"
- assert any(
- "Simulated timeout during acquire. No ledger tx was created" in m
- 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]
+ product_uuids=[product.id]
)
)
assert len(pes) == 1
assert pes[0].status == PayoutStatus.FAILED
pe = pes[0]
- # Fix the redis method
- Lock.acquire = original_acquire
-
# 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,
+ bp_pe=pe,
+ )
+ 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,
+ with pytest.raises(expected_exception=ValueError) as e:
+ pe = business_payout_event_manager.create_bp_payout_event(
+ thl_ledger_manager=thl_ledger_manager,
product=product,
created=now,
amount=rand_amount,
- payout_type=PayoutType.ACH,
+ ext_ref_id=ext_ref_id,
)
- assert e.type is ValueError
- assert "Payout event already exists!" in str(e.value)
+ assert (
+ "Cannot create a BusinessPayoutEvent with an existing transaction_id"
+ in str(e.value)
+ )
# 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)
- with pytest.raises(LedgerTransactionConditionFailedError) as e:
- pe = brokerage_product_payout_event_manager.create_bp_payout_event(
- thl_ledger_manager=thl_lm,
- product=product,
- created=now,
- amount=rand_amount,
- payout_type=PayoutType.ACH,
- )
- # But it will fail due to 1 per day check
- assert str(e.value) == ">1 tx per day"
- pe = brokerage_product_payout_event_manager.get_by_uuid(e.value.pe_uuid)
- assert pe.status == PayoutStatus.FAILED
-
- # And if we really want to, we can make it again
- now = datetime.now(tz=timezone.utc)
- pe = brokerage_product_payout_event_manager.create_bp_payout_event(
- thl_ledger_manager=thl_lm,
+ # Change the ext_ref_id so it'll create a new payout event
+ pe = business_payout_event_manager.create_bp_payout_event(
+ thl_ledger_manager=thl_ledger_manager,
product=product,
created=now,
amount=rand_amount,
- payout_type=PayoutType.ACH,
- skip_one_per_day_check=True,
- skip_wallet_balance_check=True,
+ ext_ref_id=uuid4().hex,
)
- 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
-
- Lock.release = original_release
- Lock.acquire = original_acquire
+ assert (
+ thl_ledger_manager.get_account_balance(bp_wallet_account) == 0 - rand_amount
+ )
def test_create_with_redis_error_release(
- self, product, caplog, thl_lm, brokerage_product_payout_event_manager
+ self,
+ product: Product,
+ thl_ledger_manager: ThlLedgerManager,
+ business_payout_event_manager: BusinessPayoutEventManager,
+ brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager,
+ monkeypatch: pytest.MonkeyPatch,
+ caplog: pytest.LogCaptureFixture,
):
caplog.set_level("WARNING")
- original_release = Lock.release
-
rand_amount: USDCent = USDCent(randint(100, 1_000))
- now = datetime.now(tz=timezone.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)
+ now = datetime.now(tz=UTC)
+ bp_wallet_account = thl_ledger_manager.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(
+ 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,
- product=product,
- created=now,
- amount=rand_amount,
- payout_type=PayoutType.ACH,
- )
- assert any(
- "Simulated timeout during release but ledger tx exists" in m
- for m in caplog.messages
- )
+ caplog.clear()
+ with monkeypatch.context() as m, caplog.at_level("WARNING"):
+ m.setattr(Lock, "release", broken_release)
+ business_payout_event_manager.create_bp_payout_event(
+ thl_ledger_manager=thl_ledger_manager,
+ product=product,
+ created=now,
+ amount=rand_amount,
+ ext_ref_id=uuid4().hex,
+ )
+ assert "Redis error: Simulated timeout during release" 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]
+ product_uuids=[product.uuid]
)
)
assert len(pes) == 1
assert pes[0].status == PayoutStatus.COMPLETE
- Lock.release = original_release
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..aa3b378 100644
--- a/tests/managers/thl/test_ledger/test_thl_lm_tx.py
+++ b/tests/managers/thl/test_ledger/test_thl_lm_tx.py
@@ -1,112 +1,137 @@
+from __future__ import annotations
+
import logging
-from datetime import datetime, timezone, timedelta
+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 USDCent
+from generalresearch.managers.thl.ledger_manager.exceptions import (
+ LedgerTransactionConditionFailedError,
+)
from generalresearch.managers.thl.ledger_manager.ledger import (
LedgerTransaction,
)
-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
-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,
+ Product,
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
+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")
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,
- user,
- account_revenue_task_complete,
- create_main_accounts,
- thl_lm,
- lm,
+ wall: Wall,
+ user: User,
+ account_revenue_task_complete: LedgerAccount,
+ 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,
+ 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,
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,
)
- 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,
+ 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"
@@ -115,42 +140,41 @@ 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
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(
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,
)
- 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,
+ 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,
@@ -161,14 +185,19 @@ 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],
+ bare_session: Session,
+ user: User,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
):
"""Create Wall event Complete, and Create a Tx Task Adjustment
@@ -176,29 +205,34 @@ 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_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()
- # 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)
@@ -208,28 +242,32 @@ class TestThlLedgerTxManager:
# Update the finished timestamp, but nothing else. This means that
# 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), status=Status.FAIL
)
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,
- 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,
@@ -240,7 +278,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
@@ -248,35 +286,42 @@ 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(
- product=product,
- amount=rand_amount,
- payoutevent_uuid=uuid4().hex,
- created=datetime.now(tz=timezone.utc),
- skip_wallet_balance_check=False,
- skip_one_per_day_check=False,
- skip_flag_check=False,
- )
+ with (
+ caplog.at_level(logging.INFO),
+ pytest.raises(expected_exception=LedgerTransactionConditionFailedError),
+ ):
+ thl_ledger_manager.create_tx_bp_payout(
+ product=product,
+ amount=rand_amount,
+ payoutevent_uuid=uuid4().hex,
+ created=datetime.now(tz=UTC),
+ skip_wallet_balance_check=False,
+ skip_one_per_day_check=False,
+ skip_flag_check=False,
+ )
assert "failed condition check >1 tx per day" in caplog.text
- def test_create_tx_bp_payout_(self, product, thl_lm, lm, currency):
+ def test_create_tx_bp_payout_(
+ self,
+ product: Product,
+ 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,
- created=datetime.now(tz=timezone.utc),
+ created=datetime.now(tz=UTC),
)
# Check the basic attributes
@@ -290,17 +335,21 @@ 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,
+ 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=timezone.utc),
+ created=datetime.now(tz=UTC),
direction=Direction.DEBIT,
skip_flag_check=False,
)
@@ -309,13 +358,17 @@ 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,
+ 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.
@@ -325,10 +378,10 @@ 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=timezone.utc),
+ created=datetime.now(tz=UTC),
direction=Direction.DEBIT,
)
@@ -336,32 +389,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=timezone.utc),
+ 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,
@@ -374,7 +427,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,
@@ -385,12 +438,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,
@@ -411,21 +464,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],
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
):
- delete_ledger_db()
pe = UserPayoutEvent(
uuid=uuid4().hex,
@@ -436,36 +487,32 @@ 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,
+ 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,
@@ -477,7 +524,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,
@@ -485,17 +532,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),
@@ -508,18 +557,19 @@ 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,
+ 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(
@@ -531,7 +581,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,
@@ -541,12 +591,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"),
@@ -555,19 +607,20 @@ 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,
+ 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(
@@ -579,17 +632,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,
@@ -600,19 +655,18 @@ 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,
+ 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(
@@ -624,43 +678,44 @@ 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,
+ 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,
@@ -668,44 +723,47 @@ 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,
+ 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:
@@ -713,12 +771,19 @@ class TestThlLedgerTxManagerFlows:
examples
"""
- def test_create_tx_task_complete(
- self, user, create_main_accounts, thl_lm, lm, currency, delete_ledger_db
- ):
+ @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,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ currency: LedgerCurrency,
+ ):
+
wall1 = Wall(
user_id=1,
source=Source.DYNATA,
@@ -727,10 +792,12 @@ 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_ledger_manager.create_tx_task_complete(
+ wall=wall1, user=user, created=wall1.started
)
- thl_lm.create_tx_task_complete(wall=wall1, user=user, created=wall1.started)
wall2 = Wall(
user_id=1,
@@ -740,41 +807,43 @@ 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_ledger_manager.create_tx_task_complete(
+ wall=wall2, user=user, created=wall2.started
)
- thl_lm.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,
@@ -783,7 +852,11 @@ class TestThlLedgerTxManagerFlows:
)
def test_create_transaction_task_complete_1_cent(
- self, user, create_main_accounts, thl_lm, lm, currency
+ self,
+ user: User,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ currency: LedgerCurrency,
):
wall1 = Wall(
user_id=1,
@@ -793,10 +866,10 @@ 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(
+ tx = thl_ledger_manager.create_tx_task_complete(
wall=wall1, user=user, created=wall1.started
)
@@ -804,17 +877,13 @@ 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,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ currency: LedgerCurrency,
+ session_factory: Callable[..., Session],
+ utc_hour_ago: datetime,
):
- delete_ledger_db()
- create_main_accounts()
s1: Session = session_factory(
user=user,
@@ -824,50 +893,53 @@ 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()
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)
+ 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,
+ 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)
@@ -880,11 +952,11 @@ 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(
+ tx = thl_ledger_manager.create_tx_task_complete(
wall=wall1, user=user, created=wall1.started
)
assert isinstance(tx, LedgerTransaction)
@@ -893,24 +965,27 @@ 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)
- 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,
+ user: User,
+ 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
@@ -922,34 +997,33 @@ 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)
+ 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": 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)
+ 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,
+ 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...
@@ -963,22 +1037,22 @@ 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_ledger_manager.create_tx_task_complete(
+ wall=wall1, user=user, created=wall1.started
)
- thl_lm.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": 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)
@@ -986,22 +1060,19 @@ 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,
+ 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(
@@ -1013,10 +1084,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,
@@ -1025,55 +1098,59 @@ 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:
+ @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,
- 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,
+ 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)
@@ -1089,7 +1166,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,
@@ -1102,7 +1179,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,
@@ -1110,24 +1187,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
)
@@ -1138,46 +1217,42 @@ 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,
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()
- create_main_accounts()
s1 = session_factory(
user=user,
@@ -1190,11 +1265,15 @@ 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()
- 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,
@@ -1203,21 +1282,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
@@ -1235,22 +1318,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(
@@ -1259,32 +1342,28 @@ 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),
)
- new_status, new_payout, new_user_payout = s1.determine_new_status_and_payouts()
+ _, _, _ = 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,
- 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,
+ 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
@@ -1302,51 +1381,46 @@ 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)
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": 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)
+ 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,
+ session_factory: Callable[..., Session],
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()
+
user: User = user_factory(product=product_user_wallet_no)
s1: Session = session_factory(
user=user, final_status=Status.ABANDON, wall_req_cpi=Decimal(1)
@@ -1360,9 +1434,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(
@@ -1371,24 +1445,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
)
- 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()
+ bp_commission_account = thl_ledger_manager.get_account_or_create_bp_commission(
+ product=user.product
+ )
+ 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
@@ -1399,24 +1475,20 @@ 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,
caplog,
- thl_lm,
- lm,
- currency,
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ currency: LedgerCurrency,
):
- 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
@@ -1447,7 +1519,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)
@@ -1464,7 +1536,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)
@@ -1485,34 +1557,38 @@ 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)
+ 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()
@@ -1522,7 +1598,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
)
@@ -1533,16 +1609,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(
@@ -1551,16 +1627,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(
@@ -1568,7 +1644,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 = (
@@ -1581,13 +1657,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(
@@ -1595,7 +1675,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()
)
@@ -1606,27 +1686,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,
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()
+
user: User = user_factory(product=product_user_wallet_no)
wall1 = Wall(
@@ -1640,7 +1723,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)
@@ -1656,32 +1739,34 @@ 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)
session = Session(started=wall1.started, user=user, wall_events=[wall1, wall2])
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": utc_hour_ago + timedelta(minutes=10),
- "payout": bp_pay,
- "user_payout": user_pay,
- }
- )
- thl_lm.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()
+ status=status,
+ status_code_1=status_code_1,
+ finished=utc_hour_ago + timedelta(minutes=10),
+ payout=bp_pay,
+ user_payout=user_pay,
+ )
+ thl_ledger_manager.create_tx_bp_payment(session, created=wall1.started)
+
+ revenue = thl_ledger_manager.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(
@@ -1689,17 +1774,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(
@@ -1708,14 +1793,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(
@@ -1723,13 +1808,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(
@@ -1737,13 +1826,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(
@@ -1751,12 +1844,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 1e7146a..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
@@ -1,30 +1,37 @@
+from __future__ import annotations
+
import logging
-from datetime import datetime, timezone, timedelta
+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.exceptions import (
- LedgerTransactionFlagAlreadyExistsError,
LedgerTransactionConditionFailedError,
+ LedgerTransactionFlagAlreadyExistsError,
)
-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
+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:
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()
@@ -40,16 +47,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(
@@ -62,36 +69,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,
@@ -99,12 +110,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()
@@ -118,40 +129,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_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)
+ 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()
@@ -166,15 +179,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",
@@ -182,68 +195,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_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)
- 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(timezone.utc) - timedelta(hours=1)
user: User = user_factory(product=product_amt_true)
pe = UserPayoutEvent(
@@ -254,41 +267,48 @@ 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(
- 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_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)
- 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,
@@ -297,17 +317,18 @@ 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(
- 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
@@ -315,12 +336,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()
@@ -337,64 +358,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=timezone.utc) - timedelta(hours=1)
user: User = user_factory(product=product_amt_true)
# debit_account_uuid nothing checks they match the ledger ... todo?
@@ -407,8 +429,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",
@@ -416,79 +438,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",
@@ -496,7 +530,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 5fb9e7d..fb35aa4 100644
--- a/tests/managers/thl/test_ledger/test_thl_pem.py
+++ b/tests/managers/thl/test_ledger/test_thl_pem.py
@@ -1,21 +1,39 @@
-import uuid
+from __future__ import annotations
+
+from collections.abc import Callable
from random import randint
-from uuid import uuid4, UUID
+from typing import TYPE_CHECKING
+from uuid import UUID, uuid4
import pytest
from generalresearch.currency import USDCent
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.payout import (
+ BrokerageProductPayoutEvent,
+)
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.payout import UserPayoutEvent
+ from generalresearch.models.thl.product import Product
+
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
@@ -27,35 +45,31 @@ 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,
):
- from generalresearch.models.thl.payout import UserPayoutEvent
N_PRODUCTS = randint(3, 10)
N_PAYOUT_EVENTS = randint(3, 10)
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
- )
- 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(
@@ -67,36 +81,32 @@ 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,
):
- from generalresearch.models.thl.payout import UserPayoutEvent
N_PRODUCTS = randint(3, 10)
N_PAYOUT_EVENTS = randint(3, 10)
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)
- brokerage_product_payout_event_manager.set_account_lookup_table(
- thl_lm=thl_lm
- )
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=product)
- 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]
+ product_uuids=[product.id]
)
assert len(res) == N_PAYOUT_EVENTS
@@ -105,7 +115,7 @@ 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]
+ product_uuids=[i.uuid for i in products],
)
)
@@ -113,13 +123,12 @@ 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
"""
from generalresearch.models.thl.payout import (
- UserPayoutEvent,
PayoutType,
)
@@ -145,11 +154,15 @@ class TestThlPayoutEventManager:
# def test_filter_by(self):
# raise NotImplementedError
- def test_create(self, user_payout_event_manager):
+ 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
@@ -167,27 +180,26 @@ 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,
):
- from generalresearch.models.thl.payout import UserPayoutEvent
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
+ )
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,
@@ -196,15 +208,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
@@ -212,13 +226,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()
@@ -227,10 +241,9 @@ 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)
- 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)
@@ -238,7 +251,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]
+ 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 ecf146f..6b6ef5b 100644
--- a/tests/managers/thl/test_ledger/test_user_txs.py
+++ b/tests/managers/thl/test_ledger/test_user_txs.py
@@ -1,54 +1,58 @@
-from datetime import datetime, timedelta, timezone
+from __future__ import annotations
+
+from collections.abc import Callable
+from datetime import UTC, datetime
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
from generalresearch.managers.thl.user_compensate import user_compensate
from generalresearch.models.thl.definitions import (
Status,
- WallAdjustedStatus,
)
from generalresearch.models.thl.ledger import (
TransactionType,
UserLedgerTransactionTypesSummary,
UserLedgerTransactionTypeSummary,
)
+from generalresearch.models.thl.wallet.definitions import PayoutType
if TYPE_CHECKING:
- from generalresearch.config import GRLSettings
+ 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
- from generalresearch.models.thl.wallet import PayoutType
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,
+ 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"))
- bribe_uuid = user_compensate(
- ledger_manager=thl_lm,
+ user_compensate(
+ ledger_manager=thl_ledger_manager,
user=user,
amount_int=100,
)
@@ -60,9 +64,9 @@ def test_user_txs(
amount=5,
created=utc_now,
payout_type=PayoutType.AMT_HIT,
- request_data=dict(),
+ request_data={},
)
- thl_lm.create_tx_user_payout_request(
+ thl_ledger_manager.create_tx_user_payout_request(
user=user,
payout_event=pe,
)
@@ -73,9 +77,9 @@ def test_user_txs(
amount=127,
created=utc_now,
payout_type=PayoutType.AMT_BONUS,
- request_data=dict(),
+ request_data={},
)
- thl_lm.create_tx_user_payout_request(
+ thl_ledger_manager.create_tx_user_payout_request(
user=user,
payout_event=pe,
)
@@ -92,16 +96,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}
@@ -136,33 +140,29 @@ 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",
+ ledger_manager: LedgerManager,
+ thl_ledger_manager: ThlLedgerManager,
delete_ledger_db: Callable[..., None],
- session_with_tx_factory: Callable[..., "Session"],
- adj_to_fail_with_tx_factory,
- user_payout_event_manager,
- utc_now: datetime,
):
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=thl_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
@@ -171,7 +171,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
@@ -179,7 +179,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
@@ -187,14 +187,14 @@ 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,
+ ledger_manager=thl_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
@@ -203,8 +203,8 @@ def test_user_txs_pagination(
assert txs.summary.user_bonus.entry_count == 1
# And filtering with 0 results
- now = datetime.now(tz=timezone.utc)
- txs = thl_lm.get_user_txs(user, page=1, size=5, time_start=now)
+ now = datetime.now(tz=UTC)
+ 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
@@ -215,16 +215,13 @@ def test_user_txs_pagination(
def test_user_txs_rolling_balance(
- user_factory: Callable[..., "User"],
- product_amt_true: "Product",
- create_main_accounts,
- thl_lm,
- lm,
+ user_factory: Callable[..., User],
+ product_amt_true: Product,
+ create_main_accounts: Callable[..., None],
+ thl_ledger_manager: ThlLedgerManager,
delete_ledger_db: Callable[..., None],
- session_with_tx_factory,
- adj_to_fail_with_tx_factory,
- user_payout_event_manager,
- settings: "GRLSettings",
+ user_payout_event_manager: UserPayoutEventManager,
+ settings: GRLBaseSettings,
):
"""
Creates 3 $1.00 bonuses (postive),
@@ -237,11 +234,11 @@ def test_user_txs_rolling_balance(
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)
for _ in range(3):
user_compensate(
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
user=user,
amount_int=100,
skip_flag_check=True,
@@ -253,21 +250,21 @@ def test_user_txs_rolling_balance(
cashout_method_uuid=settings.amt_bonus_cashout_method_id,
amount=150,
payout_type=PayoutType.AMT_BONUS,
- request_data=dict(),
+ request_data={},
)
- thl_lm.create_tx_user_payout_request(
+ thl_ledger_manager.create_tx_user_payout_request(
user=user,
payout_event=pe,
)
for _ in range(3):
user_compensate(
- ledger_manager=thl_lm,
+ ledger_manager=thl_ledger_manager,
user=user,
amount_int=100,
skip_flag_check=True,
)
- txs = thl_lm.get_user_txs(user, page=1, size=10)
+ txs = thl_ledger_manager.get_user_txs(user, page=1, size=10)
assert txs.transactions[0].balance_after == 100
assert txs.transactions[1].balance_after == 200
assert txs.transactions[2].balance_after == 300
@@ -278,7 +275,7 @@ def test_user_txs_rolling_balance(
# Ascending order, get 2nd page, make sure the balances include
# the previous txs. (will return last 3 txs)
- txs = thl_lm.get_user_txs(user, page=2, size=4)
+ txs = thl_ledger_manager.get_user_txs(user, page=2, size=4)
assert len(txs.transactions) == 3
assert txs.transactions[0].balance_after == 250
assert txs.transactions[1].balance_after == 350
@@ -286,7 +283,7 @@ def test_user_txs_rolling_balance(
# Descending order, get 1st page. Will
# return most recent 3 txs in desc order
- txs = thl_lm.get_user_txs(user, page=1, size=3, order_by="-created")
+ txs = thl_ledger_manager.get_user_txs(user, page=1, size=3, order_by="-created")
assert len(txs.transactions) == 3
assert txs.transactions[0].balance_after == 450
assert txs.transactions[1].balance_after == 350
diff --git a/tests/managers/thl/test_ledger/test_wallet.py b/tests/managers/thl/test_ledger/test_wallet.py
index a0abd7c..0a1da73 100644
--- a/tests/managers/thl/test_ledger/test_wallet.py
+++ b/tests/managers/thl/test_ledger/test_wallet.py
@@ -1,20 +1,31 @@
+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.models.thl.product import (
- UserWalletConfig,
PayoutConfig,
PayoutTransformation,
PayoutTransformationPercentArgs,
+ 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()
-def schrute_product(product_manager):
- 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(
@@ -27,25 +38,31 @@ 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)
+ user: User = user_factory(product=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_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 +72,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(
+ user: User = user_factory(product=schrute_product)
+ thl_ledger_manager.create_tx_user_bonus(
user=user,
amount=Decimal(1),
ref_uuid=uuid4().hex,
@@ -69,10 +90,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 c588c58..e44fe49 100644
--- a/tests/managers/thl/test_maxmind.py
+++ b/tests/managers/thl/test_maxmind.py
@@ -1,23 +1,6 @@
-import json
-import logging
-from typing 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,
-)
-from generalresearch.models.thl.maxmind.definitions import UserType
-
fake = Faker()
US_STATES = {x.lower() for x in USAddressProvider.states}
@@ -29,245 +12,244 @@ IP_v6_US = "2600:1700:ece0:9410:55d:faf3:c15d:6e4"
IP_v6_US_SAME_64 = "2600:1700:ece0:9410:55d:faf3:c15d:aaaa"
-@pytest.fixture(scope="session")
-def delete_ipinfo(thl_web_rw) -> Callable:
- def _delete_ipinfo(ip):
- thl_web_rw.execute_write(
- query="DELETE FROM thl_geoname WHERE geoname_id IN (SELECT geoname_id FROM thl_ipinformation WHERE ip = %s);",
- params=[ip],
- )
- thl_web_rw.execute_write(
- query="DELETE FROM thl_ipinformation WHERE ip = %s;",
- params=[ip],
- )
-
- return _delete_ipinfo
-
-
-class TestMaxmindBasicManager:
-
- def test_init(self, maxmind_basic_manager):
-
- assert isinstance(maxmind_basic_manager, MaxmindBasicManager)
-
- def test_get_basic_ip_information(self, maxmind_basic_manager):
- ip = IP_v4_INDIA
- maxmind_basic_manager.run_update_geoip_db()
-
- res1 = maxmind_basic_manager.get_basic_ip_information(ip_address=ip)
- assert isinstance(res1, geoip2.models.Country)
- assert res1.country.iso_code == "IN"
- assert res1.country.name == "India"
-
- res2 = maxmind_basic_manager.get_basic_ip_information(
- ip_address=fake.ipv4_private()
- )
- assert res2 is None
-
- def test_get_country_iso_from_ip_geoip2db(self, maxmind_basic_manager):
- ip = IP_v4_INDIA
- maxmind_basic_manager.run_update_geoip_db()
-
- res1 = maxmind_basic_manager.get_country_iso_from_ip_geoip2db(ip=ip)
- assert res1 == "in"
-
- res2 = maxmind_basic_manager.get_country_iso_from_ip_geoip2db(
- ip=fake.ipv4_private()
- )
- assert res2 is None
-
- def test_get_basic_ip_information_ipv6(self, maxmind_basic_manager):
- ip = IP_v6_INDIA
- maxmind_basic_manager.run_update_geoip_db()
-
- res1 = maxmind_basic_manager.get_basic_ip_information(ip_address=ip)
- assert isinstance(res1, geoip2.models.Country)
- assert res1.country.iso_code == "IN"
- assert res1.country.name == "India"
-
-
-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)
- assert isinstance(instance, MaxmindManager)
- assert isinstance(maxmind_manager, MaxmindManager)
-
- def test_create_basic(
- self,
- maxmind_manager: MaxmindManager,
- geoipinfo_manager: GeoIpInfoManager,
- delete_ipinfo,
- ):
- # This is (currently) an IP in India, and so it should only do the basic lookup
- ip = IP_v4_INDIA
- delete_ipinfo(ip)
- geoipinfo_manager.clear_cache(ip)
- assert geoipinfo_manager.get_cache(ip) is None
- assert geoipinfo_manager.get_mysql_if_exists(ip) is None
-
- maxmind_manager.run_ip_information(ip, force_insights=False)
- # Check that it is in the cache and in mysql
- res = geoipinfo_manager.get_cache(ip)
- assert res.ip == ip
- assert res.basic
- res = geoipinfo_manager.get_mysql(ip)
- assert res.ip == ip
- assert res.basic
-
- def test_create_basic_ipv6(
- self,
- maxmind_manager: MaxmindManager,
- geoipinfo_manager: GeoIpInfoManager,
- delete_ipinfo,
- ):
- # This is (currently) an IP in India, and so it should only do the basic lookup
- ip = IP_v6_INDIA
- normalized_ip, lookup_prefix = normalize_ip(ip)
- delete_ipinfo(ip)
- geoipinfo_manager.clear_cache(ip)
- delete_ipinfo(normalized_ip)
- geoipinfo_manager.clear_cache(normalized_ip)
- assert geoipinfo_manager.get_cache(ip) is None
- assert geoipinfo_manager.get_cache(normalized_ip) is None
- assert geoipinfo_manager.get_mysql_if_exists(ip) is None
- assert geoipinfo_manager.get_mysql_if_exists(normalized_ip) is None
-
- maxmind_manager.run_ip_information(ip, force_insights=False)
-
- # Check that it is in the cache
- res = geoipinfo_manager.get_cache(ip)
- # The looked up IP (/128) is returned,
- assert res.ip == ip
- assert res.lookup_prefix == "/64"
- assert res.basic
-
- # ... but the normalized version was stored (/64)
- assert geoipinfo_manager.get_cache_raw(ip) is None
- res = json.loads(geoipinfo_manager.get_cache_raw(normalized_ip))
- assert res["ip"] == normalized_ip
-
- # Check mysql
- res = geoipinfo_manager.get_mysql(ip)
- assert res.ip == ip
- assert res.lookup_prefix == "/64"
- assert res.basic
- with pytest.raises(AssertionError):
- geoipinfo_manager.get_mysql_raw(ip)
- res = geoipinfo_manager.get_mysql_raw(normalized_ip)
- assert res["ip"] == normalized_ip
-
- def test_create_insights(
- self,
- maxmind_manager: MaxmindManager,
- geoipinfo_manager: GeoIpInfoManager,
- delete_ipinfo,
- ):
- # This is (currently) an IP in the US, so it should do insights
- ip = IP_v4_US
- delete_ipinfo(ip)
- geoipinfo_manager.clear_cache(ip)
- assert geoipinfo_manager.get_cache(ip) is None
- assert geoipinfo_manager.get_mysql_if_exists(ip) is None
-
- res1 = maxmind_manager.run_ip_information(ip, force_insights=False)
- assert isinstance(res1, GeoIPInformation)
-
- # Check that it is in the cache and in mysql
- res2 = geoipinfo_manager.get_cache(ip)
- assert isinstance(res2, GeoIPInformation)
- assert res2.ip == ip
- assert not res2.basic
-
- res3 = geoipinfo_manager.get_mysql(ip)
- assert isinstance(res3, GeoIPInformation)
- assert res3.ip == ip
- assert not res3.basic
- assert res3.is_anonymous is False
- assert res3.subdivision_1_name.lower() in US_STATES
- # this might change ...
- assert res3.user_type == UserType.CELLULAR
-
- assert res1 == res2 == res3, "runner, cache, mysql all return same instance"
-
- def test_create_insights_ipv6(
- self,
- maxmind_manager: MaxmindManager,
- geoipinfo_manager: GeoIpInfoManager,
- delete_ipinfo,
- ):
- # This is (currently) an IP in the US, so it should do insights
- ip = IP_v6_US
- normalized_ip, lookup_prefix = normalize_ip(ip)
- delete_ipinfo(ip)
- geoipinfo_manager.clear_cache(ip)
- delete_ipinfo(normalized_ip)
- geoipinfo_manager.clear_cache(normalized_ip)
- assert geoipinfo_manager.get_cache(ip) is None
- assert geoipinfo_manager.get_cache(normalized_ip) is None
- assert geoipinfo_manager.get_mysql_if_exists(ip) is None
- assert geoipinfo_manager.get_mysql_if_exists(normalized_ip) is None
-
- res1 = maxmind_manager.run_ip_information(ip, force_insights=False)
- assert isinstance(res1, GeoIPInformation)
- assert res1.lookup_prefix == "/64"
-
- # Check that it is in the cache and in mysql
- res2 = geoipinfo_manager.get_cache(ip)
- assert isinstance(res2, GeoIPInformation)
- assert res2.ip == ip
- assert not res2.basic
-
- res3 = geoipinfo_manager.get_mysql(ip)
- assert isinstance(res3, GeoIPInformation)
- assert res3.ip == ip
- assert not res3.basic
- assert res3.is_anonymous is False
- assert res3.subdivision_1_name.lower() in US_STATES
- # this might change ...
- assert res3.user_type == UserType.RESIDENTIAL
-
- assert res1 == res2 == res3, "runner, cache, mysql all return same instance"
-
- def test_get_or_create_ip_information(self, maxmind_manager):
- ip = IP_v4_US
-
- res1 = maxmind_manager.get_or_create_ip_information(ip_address=ip)
- assert isinstance(res1, GeoIPInformation)
-
- res2 = maxmind_manager.get_or_create_ip_information(
- ip_address=fake.ipv4_private()
- )
- assert res2 is None
-
- def test_get_or_create_ip_information_ipv6(
- self, maxmind_manager, delete_ipinfo, geoipinfo_manager, caplog
- ):
- ip = IP_v6_US
- normalized_ip, lookup_prefix = normalize_ip(ip)
- delete_ipinfo(normalized_ip)
- geoipinfo_manager.clear_cache(normalized_ip)
-
- with caplog.at_level(logging.INFO):
- res1 = maxmind_manager.get_or_create_ip_information(ip_address=ip)
- assert isinstance(res1, GeoIPInformation)
- assert res1.ip == ip
- # It looks up in insight using the normalize IP!
- assert f"get_insights_ip_information: {normalized_ip}" in caplog.text
-
- # And it should NOT do the lookup again with an ipv6 in the same /64 block!
- ip = IP_v6_US_SAME_64
- caplog.clear()
- with caplog.at_level(logging.INFO):
- res2 = maxmind_manager.get_or_create_ip_information(ip_address=ip)
- assert isinstance(res2, GeoIPInformation)
- assert res2.ip == ip
- assert "get_insights_ip_information" not in caplog.text
-
- def test_run_ip_information(self, maxmind_manager):
- ip = IP_v4_US
-
- res = maxmind_manager.run_ip_information(ip_address=ip)
- assert isinstance(res, GeoIPInformation)
- assert res.country_name == "United States"
- assert res.country_iso == "us"
+# @pytest.fixture(scope="session")
+# def delete_ipinfo(thl_web_rw) -> Callable:
+# def _delete_ipinfo(ip):
+# thl_web_rw.execute_write(
+# query="DELETE FROM thl_geoname WHERE geoname_id IN (SELECT geoname_id FROM thl_ipinformation WHERE ip = %s);",
+# params=[ip],
+# )
+# thl_web_rw.execute_write(
+# query="DELETE FROM thl_ipinformation WHERE ip = %s;",
+# params=[ip],
+# )
+
+# return _delete_ipinfo
+
+
+# @pytest.skip("TODO: Replace with GRIP Client")
+# class TestMaxmindBasicManager:
+
+# def test_init(self,):
+
+# def test_get_basic_ip_information(self, maxmind_basic_manager):
+# ip = IP_v4_INDIA
+# maxmind_basic_manager.run_update_geoip_db()
+
+# res1 = maxmind_basic_manager.get_basic_ip_information(ip_address=ip)
+# # assert isinstance(res1, geoip2.models.Country)
+# assert res1.country.iso_code == "IN"
+# assert res1.country.name == "India"
+
+# res2 = maxmind_basic_manager.get_basic_ip_information(
+# ip_address=fake.ipv4_private()
+# )
+# assert res2 is None
+
+# def test_get_country_iso_from_ip_geoip2db(self, maxmind_basic_manager):
+# ip = IP_v4_INDIA
+# maxmind_basic_manager.run_update_geoip_db()
+
+# res1 = maxmind_basic_manager.get_country_iso_from_ip_geoip2db(ip=ip)
+# assert res1 == "in"
+
+# res2 = maxmind_basic_manager.get_country_iso_from_ip_geoip2db(
+# ip=fake.ipv4_private()
+# )
+# assert res2 is None
+
+# def test_get_basic_ip_information_ipv6(self, maxmind_basic_manager):
+# ip = IP_v6_INDIA
+# maxmind_basic_manager.run_update_geoip_db()
+
+# res1 = maxmind_basic_manager.get_basic_ip_information(ip_address=ip)
+# assert isinstance(res1, geoip2.models.Country)
+# assert res1.country.iso_code == "IN"
+# assert res1.country.name == "India"
+
+
+# class TestMaxmindManager:
+
+# 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)
+
+# def test_create_basic(
+# self,
+# maxmind_manager: MaxmindManager,
+# geoipinfo_manager: GeoIpInfoManager,
+# delete_ipinfo,
+# ):
+# # This is (currently) an IP in India, and so it should only do the basic lookup
+# ip = IP_v4_INDIA
+# delete_ipinfo(ip)
+# geoipinfo_manager.clear_cache(ip)
+# assert geoipinfo_manager.get_cache(ip) is None
+# assert geoipinfo_manager.get_mysql_if_exists(ip) is None
+
+# maxmind_manager.run_ip_information(ip, force_insights=False)
+# # Check that it is in the cache and in mysql
+# res = geoipinfo_manager.get_cache(ip)
+# assert res.ip == ip
+# assert res.basic
+# res = geoipinfo_manager.get_mysql(ip)
+# assert res.ip == ip
+# assert res.basic
+
+# def test_create_basic_ipv6(
+# self,
+# maxmind_manager: MaxmindManager,
+# geoipinfo_manager: GeoIpInfoManager,
+# delete_ipinfo,
+# ):
+# # This is (currently) an IP in India, and so it should only do the basic lookup
+# ip = IP_v6_INDIA
+# normalized_ip, lookup_prefix = normalize_ip(ip)
+# delete_ipinfo(ip)
+# geoipinfo_manager.clear_cache(ip)
+# delete_ipinfo(normalized_ip)
+# geoipinfo_manager.clear_cache(normalized_ip)
+# assert geoipinfo_manager.get_cache(ip) is None
+# assert geoipinfo_manager.get_cache(normalized_ip) is None
+# assert geoipinfo_manager.get_mysql_if_exists(ip) is None
+# assert geoipinfo_manager.get_mysql_if_exists(normalized_ip) is None
+
+# maxmind_manager.run_ip_information(ip, force_insights=False)
+
+# # Check that it is in the cache
+# res = geoipinfo_manager.get_cache(ip)
+# # The looked up IP (/128) is returned,
+# assert res.ip == ip
+# assert res.lookup_prefix == "/64"
+# assert res.basic
+
+# # ... but the normalized version was stored (/64)
+# assert geoipinfo_manager.get_cache_raw(ip) is None
+# res = json.loads(geoipinfo_manager.get_cache_raw(normalized_ip))
+# assert res["ip"] == normalized_ip
+
+# # Check mysql
+# res = geoipinfo_manager.get_mysql(ip)
+# assert res.ip == ip
+# assert res.lookup_prefix == "/64"
+# assert res.basic
+# with pytest.raises(AssertionError):
+# geoipinfo_manager.get_mysql_raw(ip)
+# res = geoipinfo_manager.get_mysql_raw(normalized_ip)
+# assert res["ip"] == normalized_ip
+
+# def test_create_insights(
+# self,
+# maxmind_manager: MaxmindManager,
+# geoipinfo_manager: GeoIpInfoManager,
+# delete_ipinfo,
+# ):
+# # This is (currently) an IP in the US, so it should do insights
+# ip = IP_v4_US
+# delete_ipinfo(ip)
+# geoipinfo_manager.clear_cache(ip)
+# assert geoipinfo_manager.get_cache(ip) is None
+# assert geoipinfo_manager.get_mysql_if_exists(ip) is None
+
+# res1 = maxmind_manager.run_ip_information(ip, force_insights=False)
+# assert isinstance(res1, GeoIPInformation)
+
+# # Check that it is in the cache and in mysql
+# res2 = geoipinfo_manager.get_cache(ip)
+# assert isinstance(res2, GeoIPInformation)
+# assert res2.ip == ip
+# assert not res2.basic
+
+# res3 = geoipinfo_manager.get_mysql(ip)
+# assert isinstance(res3, GeoIPInformation)
+# assert res3.ip == ip
+# assert not res3.basic
+# assert res3.is_anonymous is False
+# assert res3.subdivision_1_name.lower() in US_STATES
+# # this might change ...
+# assert res3.user_type == UserType.CELLULAR
+
+# assert res1 == res2 == res3, "runner, cache, mysql all return same instance"
+
+# def test_create_insights_ipv6(
+# self,
+# maxmind_manager: MaxmindManager,
+# geoipinfo_manager: GeoIpInfoManager,
+# delete_ipinfo,
+# ):
+# # This is (currently) an IP in the US, so it should do insights
+# ip = IP_v6_US
+# normalized_ip, lookup_prefix = normalize_ip(ip)
+# delete_ipinfo(ip)
+# geoipinfo_manager.clear_cache(ip)
+# delete_ipinfo(normalized_ip)
+# geoipinfo_manager.clear_cache(normalized_ip)
+# assert geoipinfo_manager.get_cache(ip) is None
+# assert geoipinfo_manager.get_cache(normalized_ip) is None
+# assert geoipinfo_manager.get_mysql_if_exists(ip) is None
+# assert geoipinfo_manager.get_mysql_if_exists(normalized_ip) is None
+
+# res1 = maxmind_manager.run_ip_information(ip, force_insights=False)
+# assert isinstance(res1, GeoIPInformation)
+# assert res1.lookup_prefix == "/64"
+
+# # Check that it is in the cache and in mysql
+# res2 = geoipinfo_manager.get_cache(ip)
+# assert isinstance(res2, GeoIPInformation)
+# assert res2.ip == ip
+# assert not res2.basic
+
+# res3 = geoipinfo_manager.get_mysql(ip)
+# assert isinstance(res3, GeoIPInformation)
+# assert res3.ip == ip
+# assert not res3.basic
+# assert res3.is_anonymous is False
+# assert res3.subdivision_1_name.lower() in US_STATES
+# # this might change ...
+# assert res3.user_type == UserType.RESIDENTIAL
+
+# assert res1 == res2 == res3, "runner, cache, mysql all return same instance"
+
+# def test_get_or_create_ip_information(self, maxmind_manager):
+# ip = IP_v4_US
+
+# res1 = maxmind_manager.get_or_create_ip_information(ip_address=ip)
+# assert isinstance(res1, GeoIPInformation)
+
+# res2 = maxmind_manager.get_or_create_ip_information(
+# ip_address=fake.ipv4_private()
+# )
+# assert res2 is None
+
+# def test_get_or_create_ip_information_ipv6(
+# self, maxmind_manager, delete_ipinfo, geoipinfo_manager, caplog
+# ):
+# ip = IP_v6_US
+# normalized_ip, lookup_prefix = normalize_ip(ip)
+# delete_ipinfo(normalized_ip)
+# geoipinfo_manager.clear_cache(normalized_ip)
+
+# with caplog.at_level(logging.INFO):
+# res1 = maxmind_manager.get_or_create_ip_information(ip_address=ip)
+# assert isinstance(res1, GeoIPInformation)
+# assert res1.ip == ip
+# # It looks up in insight using the normalize IP!
+# assert f"get_insights_ip_information: {normalized_ip}" in caplog.text
+
+# # And it should NOT do the lookup again with an ipv6 in the same /64 block!
+# ip = IP_v6_US_SAME_64
+# caplog.clear()
+# with caplog.at_level(logging.INFO):
+# res2 = maxmind_manager.get_or_create_ip_information(ip_address=ip)
+# assert isinstance(res2, GeoIPInformation)
+# assert res2.ip == ip
+# assert "get_insights_ip_information" not in caplog.text
+
+# def test_run_ip_information(self, maxmind_manager):
+# ip = IP_v4_US
+
+# res = maxmind_manager.run_ip_information(ip_address=ip)
+# assert isinstance(res, GeoIPInformation)
+# assert res.country_name == "United States"
+# assert res.country_iso == "us"
diff --git a/tests/managers/thl/test_payout.py b/tests/managers/thl/test_payout.py
index 31087b8..52bbbec 100644
--- a/tests/managers/thl/test_payout.py
+++ b/tests/managers/thl/test_payout.py
@@ -1,25 +1,52 @@
+import io
import logging
import os
-from datetime import datetime, timezone, timedelta
+from collections.abc import Callable
+from datetime import UTC, datetime, timedelta
from decimal import Decimal
-from random import choice as rand_choice, randint
-from typing import Optional
+from random import choice as rand_choice
+from random import randint
+from typing import TYPE_CHECKING
from uuid import uuid4
import pandas as pd
import pytest
+from dask.distributed import Client as DaskClient
from generalresearch.currency import USDCent
-from generalresearch.managers.thl.ledger_manager.exceptions import (
- LedgerTransactionConditionFailedError,
-)
-from generalresearch.managers.thl.payout import UserPayoutEventManager
from generalresearch.models.thl.definitions import PayoutStatus
-from generalresearch.models.thl.ledger import LedgerEntry, Direction
-from generalresearch.models.thl.payout import BusinessPayoutEvent
-from generalresearch.models.thl.payout import UserPayoutEvent
-from generalresearch.models.thl.wallet import PayoutType
-from generalresearch.models.thl.ledger import LedgerAccount
+from generalresearch.models.thl.finance import BusinessBalances
+from generalresearch.models.thl.payout import (
+ BrokerageProductPayoutEvent,
+ BusinessPayoutEvent,
+)
+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()
@@ -27,17 +54,15 @@ cashout_method_uuid = uuid4().hex
class TestPayout:
-
def test_get_by_uuid_and_create(
self,
- user,
+ user: User,
user_payout_event_manager: UserPayoutEventManager,
- thl_lm,
- utc_now,
+ thl_ledger_manager: ThlLedgerManager,
+ utc_now: datetime,
):
-
- user_account: LedgerAccount = thl_lm.get_account_or_create_user_wallet(
- user=user
+ user_account: LedgerAccount = (
+ thl_ledger_manager.get_account_or_create_user_wallet(user=user)
)
pe1: UserPayoutEvent = user_payout_event_manager.create(
@@ -57,11 +82,14 @@ class TestPayout:
assert pe1 == pe2
- def test_update(self, user, user_payout_event_manager, lm, thl_lm, utc_now):
- from generalresearch.models.thl.definitions import PayoutStatus
- from generalresearch.models.thl.wallet import PayoutType
-
- user_account = thl_lm.get_account_or_create_user_wallet(user=user)
+ def test_update(
+ self,
+ user: User,
+ user_payout_event_manager: UserPayoutEventManager,
+ thl_ledger_manager: ThlLedgerManager,
+ utc_now: datetime,
+ ):
+ user_account = thl_ledger_manager.get_account_or_create_user_wallet(user=user)
pe1 = user_payout_event_manager.create(
status=PayoutStatus.PENDING,
@@ -89,142 +117,84 @@ class TestPayout:
def test_create_bp_payout(
self,
- user,
- thl_web_rr,
- user_payout_event_manager,
- lm,
- thl_lm,
- product,
- brokerage_product_payout_event_manager,
- delete_ledger_db,
- create_main_accounts,
):
- delete_ledger_db()
- create_main_accounts()
- from generalresearch.models.thl.ledger import LedgerAccount
-
- thl_lm.get_account_or_create_bp_wallet(product=product)
- brokerage_product_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
-
- with pytest.raises(expected_exception=LedgerTransactionConditionFailedError):
- # wallet balance failure
- brokerage_product_payout_event_manager.create_bp_payout_event(
- thl_ledger_manager=thl_lm,
- product=product,
- amount=USDCent(100),
- skip_wallet_balance_check=False,
- skip_one_per_day_check=False,
- )
-
- # (we don't have a special method for this) Put money in the BP's account
- amount_cents = 100
- cash_account: LedgerAccount = thl_lm.get_account_cash()
- bp_wallet: LedgerAccount = thl_lm.get_account_or_create_bp_wallet(
- product=product
- )
-
- entries = [
- LedgerEntry(
- direction=Direction.DEBIT,
- account_uuid=cash_account.uuid,
- amount=amount_cents,
- ),
- LedgerEntry(
- direction=Direction.CREDIT,
- account_uuid=bp_wallet.uuid,
- amount=amount_cents,
- ),
- ]
-
- lm.create_tx(entries=entries)
- assert 100 == lm.get_account_balance(account=bp_wallet)
-
- # Then run it again for $1.00
- brokerage_product_payout_event_manager.create_bp_payout_event(
- thl_ledger_manager=thl_lm,
- product=product,
- amount=USDCent(100),
- skip_wallet_balance_check=False,
- skip_one_per_day_check=False,
- )
- assert 0 == lm.get_account_balance(account=bp_wallet)
-
- # Run again should without balance check, should still fail due to day check
- with pytest.raises(LedgerTransactionConditionFailedError):
- brokerage_product_payout_event_manager.create_bp_payout_event(
- thl_ledger_manager=thl_lm,
- product=product,
- amount=USDCent(100),
- skip_wallet_balance_check=True,
- skip_one_per_day_check=False,
- )
+ # create_bp_payout_event does not get called directly. We have tests
+ # for the ledger methods already
+ pass
- # And then we can run again skip both checks
- pe = brokerage_product_payout_event_manager.create_bp_payout_event(
- thl_ledger_manager=thl_lm,
- product=product,
+ @pytest.fixture
+ def pending_bp_pe(
+ self,
+ thl_web_rw: PostgresConfig,
+ product: Product,
+ thl_ledger_manager: ThlLedgerManager,
+ brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager,
+ utc_now: datetime,
+ ) -> BrokerageProductPayoutEvent:
+ account = thl_ledger_manager.get_account_or_create_bp_wallet(product=product)
+ bp_pe = BrokerageProductPayoutEvent(
+ product_id=product.uuid,
amount=USDCent(100),
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
- )
- assert -100 == lm.get_account_balance(account=bp_wallet)
-
- pe = brokerage_product_payout_event_manager.get_by_uuid(pe.uuid)
- txs = lm.get_tx_filtered_by_metadata(
- metadata_key="event_payout", metadata_value=pe.uuid
+ payout_type=PayoutType.ACH,
+ debit_account_uuid=account.uuid,
+ cashout_method_uuid=brokerage_product_payout_event_manager.CASHOUT_METHOD_UUID,
+ created=utc_now,
)
-
- assert 1 == len(txs)
+ params = bp_pe.model_dump_postgres()
+ # This shouldn't exist. For testing only, so no supplier_payout
+ params["supplier_payout_id"] = None
+ thl_web_rw.execute_write(
+ """
+ INSERT INTO event_payout (uuid, debit_account_uuid, created, cashout_method_uuid,
+ amount, status, ext_ref_id, payout_type, order_data,
+ request_data, supplier_payout_id)
+ VALUES (%(uuid)s, %(debit_account_uuid)s, %(created)s, %(cashout_method_uuid)s,
+ %(amount)s, %(status)s, %(ext_ref_id)s, %(payout_type)s, %(order_data)s,
+ %(request_data)s, %(supplier_payout_id)s);
+ """,
+ params,
+ )
+ return bp_pe
def test_create_bp_payout_quick_dupe(
self,
- user,
- product,
- thl_web_rw,
- brokerage_product_payout_event_manager,
- thl_lm,
- lm,
- utc_now,
- create_main_accounts,
+ product: Product,
+ brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager,
+ thl_ledger_manager: ThlLedgerManager,
+ utc_now: datetime,
+ pending_bp_pe: BrokerageProductPayoutEvent,
):
- 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.create_bp_payout_event(
- thl_ledger_manager=thl_lm,
+ brokerage_product_payout_event_manager.create_tx_bp_payout_from_payout_event(
+ thl_ledger_manager=thl_ledger_manager,
+ bp_pe=pending_bp_pe,
product=product,
- amount=USDCent(100),
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
created=utc_now,
)
with pytest.raises(ValueError) as cm:
- brokerage_product_payout_event_manager.create_bp_payout_event(
- thl_ledger_manager=thl_lm,
+ brokerage_product_payout_event_manager.create_tx_bp_payout_from_payout_event(
+ thl_ledger_manager=thl_ledger_manager,
product=product,
- amount=USDCent(100),
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
+ bp_pe=pending_bp_pe,
created=utc_now,
)
assert "Payout event already exists!" in str(cm.value)
def test_filter(
self,
- thl_web_rw,
- thl_lm,
- lm,
- product,
- user,
- user_payout_event_manager,
- utc_now,
+ thl_ledger_manager: ThlLedgerManager,
+ product: Product,
+ user: User,
+ user_payout_event_manager: UserPayoutEventManager,
+ 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_lm.get_account_or_create_user_wallet(user=user)
- bp_account = thl_lm.get_account_or_create_bp_wallet(product=product)
+ 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)
user_payout_event_manager.create(
status=PayoutStatus.PENDING,
@@ -289,166 +259,109 @@ class TestPayout:
class TestPayoutEventManager:
-
- def test_set_account_lookup_table(
- self, payout_event_manager, thl_redis_config, thl_lm, delete_ledger_db
- ):
- delete_ledger_db()
- rc = thl_redis_config.create_redis_client()
- rc.delete("pem:account_to_product")
- rc.delete("pem:product_to_account")
- N = 5
-
- for idx in range(N):
- thl_lm.get_account_or_create_bp_wallet_by_uuid(product_uuid=uuid4().hex)
-
- res = rc.hgetall(name="pem:account_to_product")
- assert len(res.items()) == 0
-
- res = rc.hgetall(name="pem:product_to_account")
- assert len(res.items()) == 0
-
- payout_event_manager.set_account_lookup_table(
- thl_lm=thl_lm,
- )
-
- res = rc.hgetall(name="pem:account_to_product")
- assert len(res.items()) == N
-
- res = rc.hgetall(name="pem:product_to_account")
- assert len(res.items()) == N
-
- thl_lm.get_account_or_create_bp_wallet_by_uuid(product_uuid=uuid4().hex)
- payout_event_manager.set_account_lookup_table(
- thl_lm=thl_lm,
- )
-
- res = rc.hgetall(name="pem:account_to_product")
- assert len(res.items()) == N + 1
-
- res = rc.hgetall(name="pem:product_to_account")
- assert len(res.items()) == N + 1
+ pass
class TestBusinessPayoutEventManager:
-
@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 "5d"
@pytest.fixture
- def duration(self) -> Optional["timedelta"]:
+ def duration(self) -> timedelta | None:
return timedelta(days=10)
def test_base(
self,
- brokerage_product_payout_event_manager,
- business_payout_event_manager,
- delete_ledger_db,
- create_main_accounts,
- thl_lm,
- thl_web_rr,
- product_factory,
- bp_payout_factory,
- business,
+ business_payout_event_manager: BusinessPayoutEventManager,
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ thl_ledger_manager: ThlLedgerManager,
+ product_factory: Callable[..., Product],
+ brokerage_product_payout_event_factory: Callable[
+ ..., BrokerageProductPayoutEvent
+ ],
+ gr_business: Business,
):
delete_ledger_db()
create_main_accounts()
- from generalresearch.models.thl.product import Product
-
- p1: Product = product_factory(business=business)
- thl_lm.get_account_or_create_bp_wallet(product=p1)
- business_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
+ p1: Product = product_factory(business=gr_business)
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=p1)
ach_id1 = uuid4().hex
ach_id2 = uuid4().hex
- bp_payout_factory(
- product=p1,
- amount=USDCent(1),
- ext_ref_id=None,
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
+ # ext_ref_id is required now
+ 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,
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
+ 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",
+ ):
+ brokerage_product_payout_event_factory(
+ product=p1, amount=USDCent(25), ext_ref_id=ach_id1
+ )
- bp_payout_factory(
- product=p1,
- amount=USDCent(25),
- ext_ref_id=ach_id1,
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
+ brokerage_product_payout_event_factory(
+ product=p1, amount=USDCent(50), ext_ref_id=ach_id2
)
- bp_payout_factory(
- product=p1,
- amount=USDCent(50),
- ext_ref_id=ach_id2,
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
+ gr_business.prebuild_payouts(
+ bpem=business_payout_event_manager,
)
- business.prebuild_payouts(
- thl_pg_config=thl_web_rr,
- thl_lm=thl_lm,
- bpem=business_payout_event_manager,
+ 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
- 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 len(business.payouts[1].bp_payouts) == 2
+ # Cannot pay out the same product twice in the same business payout
+ # assert len(business.payouts[1].bp_payouts) == 2
+ 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 is 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,
- brokerage_product_payout_event_manager,
- business_payout_event_manager,
- delete_ledger_db,
- create_main_accounts,
- thl_lm,
- thl_web_rr,
- product_factory,
- bp_payout_factory,
- delete_df_collection,
- user_factory,
- ledger_collection,
- session_with_tx_factory,
- pop_ledger_merge,
- client_no_amm,
- mnt_filepath,
- lm,
- product_manager,
- start,
- business,
+ business_payout_event_manager: BusinessPayoutEventManager,
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ thl_ledger_manager: ThlLedgerManager,
+ thl_web_rr: PostgresConfig,
+ product_factory: Callable[..., Product],
+ 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,
+ gr_business: Business,
):
delete_ledger_db()
create_main_accounts()
delete_df_collection(coll=ledger_collection)
- from generalresearch.models.thl.product import Product
- from generalresearch.models.thl.user import User
-
- p1: Product = product_factory(business=business)
+ p1: Product = product_factory(business=gr_business)
u1: User = user_factory(product=p1)
- thl_lm.get_account_or_create_bp_wallet(product=p1)
- business_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=p1)
# $250.00 to work with
for idx in range(1, 10):
@@ -461,32 +374,34 @@ class TestBusinessPayoutEventManager:
ach_id1 = uuid4().hex
ach_id2 = uuid4().hex
- with pytest.raises(expected_exception=Warning) as cm:
+ with pytest.raises(
+ expected_exception=AssertionError, match="No Business Payout found"
+ ):
business_payout_event_manager.update_ext_reference_ids(
new_value=ach_id2,
current_value=ach_id1,
)
- assert "No event_payouts found to UPDATE" in str(cm)
# 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=lm,
+ lm=thl_ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
)
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_lm,
+ thl_lm=thl_ledger_manager,
transaction_id=ach_id1,
)
assert isinstance(res, BusinessPayoutEvent)
+ assert business_payout_event_manager.get_by_ext_ref_id(ext_ref_id=ach_id1)
# Okay, now that there is a payout_event, let's try to update the
# ext_reference_id
@@ -495,111 +410,17 @@ class TestBusinessPayoutEventManager:
current_value=ach_id1,
)
- res = business_payout_event_manager.filter_by(ext_ref_id=ach_id1)
- assert len(res) == 0
+ with pytest.raises(
+ expected_exception=AssertionError, match="No Business Payout found"
+ ):
+ business_payout_event_manager.get_by_ext_ref_id(ext_ref_id=ach_id1)
- res = business_payout_event_manager.filter_by(ext_ref_id=ach_id2)
- assert len(res) == 1
+ assert business_payout_event_manager.get_by_ext_ref_id(ext_ref_id=ach_id2)
- def test_delete_failed_business_payout(
- self,
- brokerage_product_payout_event_manager,
- business_payout_event_manager,
- delete_ledger_db,
- create_main_accounts,
- thl_lm,
- thl_web_rr,
- product_factory,
- bp_payout_factory,
- currency,
- delete_df_collection,
- user_factory,
- ledger_collection,
- session_with_tx_factory,
- pop_ledger_merge,
- client_no_amm,
- mnt_filepath,
- lm,
- product_manager,
- start,
- business,
+ def test_recoup_empty(
+ self, business_payout_event_manager: BusinessPayoutEventManager
):
- delete_ledger_db()
- create_main_accounts()
- delete_df_collection(coll=ledger_collection)
-
- from generalresearch.models.thl.product import Product
- from generalresearch.models.thl.user import User
-
- p1: Product = product_factory(business=business)
- u1: User = user_factory(product=p1)
- thl_lm.get_account_or_create_bp_wallet(product=p1)
- business_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
-
- # $250.00 to work with
- for idx in range(1, 10):
- session_with_tx_factory(
- user=u1,
- wall_req_cpi=Decimal("25.00"),
- started=start + timedelta(days=1, minutes=idx),
- )
-
- # 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(
- thl_pg_config=thl_web_rr,
- lm=lm,
- ds=mnt_filepath,
- client=client_no_amm,
- pop_ledger=pop_ledger_merge,
- )
-
- ach_id1 = uuid4().hex
-
- res = business_payout_event_manager.create_from_ach_or_wire(
- business=business,
- amount=USDCent(100_01),
- pm=product_manager,
- thl_lm=thl_lm,
- transaction_id=ach_id1,
- )
- assert isinstance(res, BusinessPayoutEvent)
-
- # (1) Confirm the initial Event Payout, Tx, TxMeta, TxEntry all exist
- event_payouts = business_payout_event_manager.filter_by(ext_ref_id=ach_id1)
- event_payout_uuids = [i.uuid for i in event_payouts]
- assert len(event_payout_uuids) == 1
- tags = [f"{currency.value}:bp_payout:{x}" for x in event_payout_uuids]
- transactions = thl_lm.get_txs_by_tags(tags=tags)
- assert len(transactions) == 1
- tx_metadata_ids = thl_lm.get_tx_metadata_ids_by_txs(transactions=transactions)
- assert len(tx_metadata_ids) == 2
- tx_entries = thl_lm.get_tx_entries_by_txs(transactions=transactions)
- assert len(tx_entries) == 2
-
- # (2) Delete!
- business_payout_event_manager.delete_failed_business_payout(
- ext_ref_id=ach_id1, thl_lm=thl_lm
- )
-
- # (3) Confirm the initial Event Payout, Tx, TxMeta, TxEntry have
- # all been deleted
- res = business_payout_event_manager.filter_by(ext_ref_id=ach_id1)
- assert len(res) == 0
-
- # Note: b/c the event_payout shouldn't exist anymore, we are taking
- # the tag strings and transactions from when they did..
- res = thl_lm.get_txs_by_tags(tags=tags)
- assert len(res) == 0
-
- tx_metadata_ids = thl_lm.get_tx_metadata_ids_by_txs(transactions=transactions)
- assert len(tx_metadata_ids) == 0
- tx_entries = thl_lm.get_tx_entries_by_txs(transactions=transactions)
- assert len(tx_entries) == 0
-
- def test_recoup_empty(self, business_payout_event_manager):
- res = {uuid4().hex: USDCent(0) for i in range(100)}
+ 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"]
@@ -609,10 +430,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"]
@@ -624,10 +447,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"]
@@ -643,7 +466,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"
@@ -657,9 +482,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)
@@ -670,14 +497,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())
@@ -697,7 +526,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
)
@@ -707,8 +536,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):
- import io
+ def test_distribute_amount(
+ self, business_payout_event_manager: BusinessPayoutEventManager
+ ):
df = pd.read_csv(
io.StringIO(
@@ -724,31 +554,31 @@ 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,
+ gr_business: Business,
+ user_factory: Callable[..., User],
+ product_factory: Callable[..., Product],
+ session_with_tx_factory: Callable[..., Session],
+ pop_ledger_merge: PopLedgerMerge,
+ start: datetime,
+ brokerage_product_payout_event_factory: Callable[
+ ..., BrokerageProductPayoutEvent
+ ],
+ ledger_manager: LedgerManager,
+ product_manager: ProductManager,
):
- """Test having a Business with three products.. one that lost money
+ """Test having a Business with three products. One that lost money
and two that gained money. Ensure that the Business balance
reflects that to compensate for the Product in the negative and only
assigns Brokerage Product payments from the 2 accounts that have
@@ -759,12 +589,9 @@ class TestBusinessPayoutEventManager:
create_main_accounts()
delete_df_collection(coll=ledger_collection)
- from generalresearch.models.thl.product import Product
- from generalresearch.models.thl.user import User
-
- p1: Product = product_factory(business=business)
+ p1: Product = product_factory(business=gr_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,
@@ -776,8 +603,7 @@ class TestBusinessPayoutEventManager:
wall_req_cpi=Decimal("5.00"),
started=start + timedelta(days=6),
)
- payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
- 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),
@@ -785,9 +611,9 @@ 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(
- thl_pg_config=thl_web_rr,
- lm=lm,
+ gr_business.prebuild_balance(
+ product_manager=product_manager,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
@@ -795,40 +621,151 @@ 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_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,
+ 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,
+ gr_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"""
+ # Now let's load it up and actually test some things
+ delete_ledger_db()
+ create_main_accounts()
+ delete_df_collection(coll=ledger_collection)
+
+ 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)
+ 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
+
+ # Product 1: Complete $10 x 20
+ for idx in range(20):
+ session_with_tx_factory(
+ user=u2,
+ wall_req_cpi=Decimal("10.00"),
+ started=start + timedelta(days=1, hours=2, minutes=1 + idx),
+ )
+
+ # Product 2: Complete $10 x 30
+ for idx in range(30):
+ session_with_tx_factory(
+ user=u3,
+ wall_req_cpi=Decimal("10.00"),
+ started=start + timedelta(days=1, hours=3, minutes=1 + idx),
+ )
+
+ 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,
+ lm=ledger_manager,
+ ds=mnt_filepath,
+ client=client_no_amm,
+ pop_ledger=pop_ledger_merge,
+ )
+
+ 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=gr_business,
+ amount=USDCent(100_00),
+ pm=product_manager,
+ thl_lm=thl_ledger_manager,
+ created=start + timedelta(days=1, hours=5),
+ transaction_id=ach_id1,
+ )
+ print(f"{bp1=}")
+ assert isinstance(bp1, BusinessPayoutEvent)
+ assert len(bp1.bp_payouts) == 2
+
+ bp2 = business_payout_event_manager.create_from_ach_or_wire(
+ business=gr_business,
+ amount=USDCent(bb.available_balance),
+ pm=product_manager,
+ thl_lm=thl_ledger_manager,
+ created=start + timedelta(days=2, hours=5),
+ transaction_id=ach_id2,
+ )
+ print(f"{bp2=}")
+ assert isinstance(bp2, BusinessPayoutEvent)
+ assert len(bp2.bp_payouts) == 2
+
+ with caplog.at_level(logging.WARNING):
+ business_payout_event_manager.resume_failed_business_payout(
+ ext_ref_id=ach_id1, thl_lm=thl_ledger_manager, pm=product_manager
+ )
+ assert "Nothing to do!" in caplog.text
+
+ # bpe = business_payout_event_manager.get_by_ext_ref_id(ext_ref_id=ach_id1)
+ # bp_pe = bpe.bp_payouts[0]
+ # thl_web_rr.execute_write(
+ # """
+ # UPDATE event_payout
+ # SET status = %(status)s
+ # WHERE uuid = %(uuid)s""",
+ # {"uuid": bp_pe.uuid, "status": PayoutStatus.FAILED},
+ # )
+
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,
+ mnt_filepath: GRLDatasets,
+ thl_ledger_manager: ThlLedgerManager,
+ client_no_amm: DaskClient,
+ 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,
+ gr_business: Business,
+ user_factory: Callable[..., User],
+ product_factory: Callable[..., Product],
+ session_with_tx_factory: Callable[..., None],
+ pop_ledger_merge: PopLedgerMerge,
+ start: datetime,
+ brokerage_product_payout_event_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
@@ -841,21 +778,17 @@ class TestBusinessPayoutEventManager:
create_main_accounts()
delete_df_collection(coll=ledger_collection)
- from generalresearch.models.thl.product import Product
- from generalresearch.models.thl.user import User
-
- 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)
- 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
# Product 1: Complete, Payout, Recon..
s1 = session_with_tx_factory(
@@ -863,14 +796,11 @@ class TestBusinessPayoutEventManager:
wall_req_cpi=Decimal("5.00"),
started=start + timedelta(days=1),
)
- payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
- bp_payout_factory(
+ brokerage_product_payout_event_factory(
product=u1.product,
amount=USDCent(475), # 95% of $5.00
ext_ref_id=ach_id1,
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,
@@ -893,18 +823,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 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(
- thl_pg_config=thl_web_rr,
- lm=lm,
+ gr_business.prebuild_balance(
+ product_manager=product_manager,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
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]
pb3 = bb1.product_balances[2]
@@ -927,20 +858,21 @@ 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_lm,
+ thl_lm=thl_ledger_manager,
bpem=business_payout_event_manager,
)
- 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_lm,
+ thl_lm=thl_ledger_manager,
created=start + timedelta(days=1, hours=5),
)
assert isinstance(bp1, BusinessPayoutEvent)
@@ -948,7 +880,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,
@@ -956,14 +888,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)
@@ -971,16 +903,15 @@ 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,
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
@@ -989,6 +920,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
@@ -1005,34 +938,29 @@ class TestBusinessPayoutEventManager:
def test_ach_payment_partial_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,
- rm_ledger_collection,
- rm_pop_ledger_merge,
+ mnt_filepath: GRLDatasets,
+ thl_ledger_manager: ThlLedgerManager,
+ client_no_amm: DaskClient,
+ payout_event_manager: PayoutEventManager,
+ business_payout_event_manager: BusinessPayoutEventManager,
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ delete_df_collection: Callable[..., None],
+ ledger_collection: LedgerDFCollection,
+ gr_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,
+ 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
+ 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
@@ -1047,18 +975,15 @@ class TestBusinessPayoutEventManager:
create_main_accounts()
delete_df_collection(coll=ledger_collection)
- from generalresearch.models.thl.product import Product
- from generalresearch.models.thl.user import User
-
- 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)
- 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):
@@ -1068,27 +993,27 @@ 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(
+ gr_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,
+ 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
assert bb1.adjustment == 0
assert bb1.payout == bb1.net
@@ -1100,24 +1025,25 @@ 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 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_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 gr_business: Business, let's confirm the updated
# balances. Clear and rebuild the parquet files.
rm_ledger_collection()
rm_pop_ledger_merge()
@@ -1127,51 +1053,46 @@ 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=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,
+ gr_business.prebuild_payouts(
bpem=business_payout_event_manager,
)
- assert len(business.payouts) == 1
- assert len(business.payouts[0].bp_payouts) == 3
- assert business.payouts_total_str == "$250.00"
- 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,
- 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,
+ payout_event_manager: PayoutEventManager,
+ business_payout_event_manager: BusinessPayoutEventManager,
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ delete_df_collection: Callable[..., None],
+ ledger_collection: LedgerDFCollection,
+ gr_business: Business,
+ user_factory: Callable[..., User],
+ product_factory: Callable[..., Product],
+ session_with_tx_factory: Callable[..., Session],
+ pop_ledger_merge: PopLedgerMerge,
+ start: datetime,
+ 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
@@ -1179,18 +1100,15 @@ class TestBusinessPayoutEventManager:
create_main_accounts()
delete_df_collection(coll=ledger_collection)
- from generalresearch.models.thl.product import Product
- from generalresearch.models.thl.user import User
-
- 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)
- 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
@@ -1202,26 +1120,26 @@ class TestBusinessPayoutEventManager:
wall_req_cpi=Decimal("7.50"),
started=start + timedelta(days=1, hours=1 + iidx, 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)
rm_ledger_collection()
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=lm,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
)
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,
- thl_lm=thl_lm,
+ thl_lm=thl_ledger_manager,
created=start + timedelta(days=2, hours=1),
)
@@ -1229,20 +1147,20 @@ 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=lm,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
)
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,
- thl_lm=thl_lm,
+ thl_lm=thl_ledger_manager,
created=start + timedelta(days=4, hours=1),
)
@@ -1253,17 +1171,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_lm,
+ thl_lm=thl_ledger_manager,
bpem=business_payout_event_manager,
)
- business.prebuild_balance(
+ gr_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 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_product.py b/tests/managers/thl/test_product.py
index 78d5dde..8d72fa5 100644
--- a/tests/managers/thl/test_product.py
+++ b/tests/managers/thl/test_product.py
@@ -1,24 +1,35 @@
+from __future__ import annotations
+
+from collections.abc import Callable
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
-from generalresearch.models import Source
+from generalresearch.models.definitions 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
+
+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):
- 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]}",
@@ -37,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):
+ def test_get_by_uuids(
+ self, product_factory: Callable[..., Product], 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_factory(
product_id=product_id,
team_id=uuid4().hex,
name=f"Test Product ID #{uuid4().hex[:6]}",
@@ -62,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):
- 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]}",
@@ -74,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):
+ 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]}",
@@ -106,13 +123,15 @@ 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_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,
@@ -124,8 +143,10 @@ class TestProductManagerGetMethods:
class TestProductManagerCreation:
- def test_base(self, product_manager):
- 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]}",
@@ -136,7 +157,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
@@ -179,20 +200,26 @@ class TestProductManager:
]
]
- def test_get_by_uuid1(self, product_manager, team, product, product_factory):
- p1 = product_factory(team=team)
+ def test_get_by_uuid1(
+ self,
+ product_manager: ProductManager,
+ gr_team: Team,
+ product: Product,
+ product_factory: Callable[..., Product],
+ ):
+ p1 = product_factory(team=gr_team)
instance = product_manager.get_by_uuid(product_uuid=p1.uuid)
assert instance.id == p1.id
# No Team and no user_create_config
- assert instance.team_id == team.uuid
+ assert instance.team_id == gr_team.uuid
# user_create_config can't be None, so ensure the default was set.
assert isinstance(instance.user_create_config, UserCreateConfig)
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
@@ -204,7 +231,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, 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
@@ -220,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):
+ 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)
@@ -235,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):
+ def test_global_sources(
+ self, product_factory: Callable[..., Product], product_manager: ProductManager
+ ):
sources_config = SupplyConfig(
policies=[
SupplyPolicy(
@@ -246,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
@@ -262,8 +295,10 @@ class TestProductManager:
p2 = product_manager.get_by_uuid(p1.id)
assert p1 == p2
- def test_user_health_config(self, product_manager):
- 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"])
)
@@ -273,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):
- 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
@@ -320,8 +355,10 @@ class TestProductManager:
class TestProductManagerUpdate:
- def test_update(self, product_manager):
- 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)
@@ -341,8 +378,10 @@ class TestProductManagerUpdate:
class TestProductManagerCacheClear:
- def test_cache_clear(self, product_manager):
- 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_product_prod.py b/tests/managers/thl/test_product_prod.py
index 7b4f677..d584527 100644
--- a/tests/managers/thl/test_product_prod.py
+++ b/tests/managers/thl/test_product_prod.py
@@ -1,17 +1,25 @@
+from __future__ import annotations
+
import logging
+from collections.abc import Callable
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
from generalresearch.models.thl.product import Product
-from test_utils.models.conftest import product_factory
+
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.product import ProductManager
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: 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)
@@ -23,7 +31,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, 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 +53,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: ProductManager
+ ):
products = [product_factory(), product_factory(), product_factory()]
instance = product_manager.get_by_uuid_if_exists(product_uuid=products[0].id)
@@ -52,7 +64,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: Callable[..., Product]
+ ):
products = [product_factory(), product_factory(), product_factory()]
res = product_manager.get_by_uuids_if_exists(
@@ -75,8 +89,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
- pass
diff --git a/tests/managers/thl/test_profiling/test_question.py b/tests/managers/thl/test_profiling/test_question.py
index 998466e..e4afb87 100644
--- a/tests/managers/thl/test_profiling/test_question.py
+++ b/tests/managers/thl/test_profiling/test_question.py
@@ -1,12 +1,20 @@
+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:
- 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 +25,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 +54,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..feab902 100644
--- a/tests/managers/thl/test_profiling/test_schema.py
+++ b/tests/managers/thl/test_profiling/test_schema.py
@@ -1,9 +1,21 @@
+from collections.abc import Callable
+from typing import TYPE_CHECKING
+
from generalresearch.models.thl.profiling.upk_property import PropertyType
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.profiling.schema import (
+ UpkSchemaManager,
+ )
+
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 +47,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 53bb8fe..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,12 @@
-from datetime import datetime, timezone
+from collections.abc import Callable
+from datetime import UTC, datetime
+from typing import TYPE_CHECKING
-from generalresearch.managers.thl.profiling.user_upk import UserUpkManager
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.profiling.user_upk import UserUpkManager
+ from generalresearch.models.thl.user import User
-now = datetime.now(tz=timezone.utc)
+now = datetime.now(tz=UTC)
base = {
"country_iso": "us",
"language_iso": "eng",
@@ -21,11 +25,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 6bedc2b..60edcb9 100644
--- a/tests/managers/thl/test_session_manager.py
+++ b/tests/managers/thl/test_session_manager.py
@@ -1,28 +1,42 @@
-from datetime import timedelta
+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.models import DeviceType
+from generalresearch.models.definitions 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
+from generalresearch.models.thl.session import Session
+from generalresearch.models.thl.user import User
+
+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()
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),
- user_payout_min=Decimal("1"),
- user_payout_max=Decimal("2"),
+ user_payout_min=Decimal(1),
+ user_payout_max=Decimal(2),
)
s1 = session_manager.create(
@@ -40,7 +54,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
@@ -60,7 +76,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)
@@ -68,7 +84,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)
@@ -76,14 +94,16 @@ 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: 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)
@@ -96,42 +116,40 @@ class TestSessionManagerFilter:
def test_team(
self,
- product_factory,
- user_factory,
- team,
- session_manager,
- user,
- utc_hour_ago,
- thl_web_rr,
+ product_factory: Callable[..., Product],
+ user_factory: Callable[..., User],
+ gr_team: Team,
+ session_manager: SessionManager,
+ utc_hour_ago: datetime,
+ thl_web_rr: PostgresConfig,
):
- p1 = product_factory(team=team)
+ p1 = product_factory(team=gr_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)
- team.prefetch_products(thl_pg_config=thl_web_rr)
- assert len(team.product_uuids) == 1
- res = session_manager.filter(product_uuids=team.product_uuids)
+ gr_team.prefetch_products(thl_pg_config=thl_web_rr)
+ assert len(gr_team.product_uuids) == 1
+ res = session_manager.filter(product_uuids=gr_team.product_uuids)
assert len(res) == 5
def test_business(
self,
- product_factory,
- business,
- user_factory,
- session_manager,
- user,
- utc_hour_ago,
- thl_web_rr,
+ product_factory: Callable[..., Product],
+ 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 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)
- 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/managers/thl/test_survey.py b/tests/managers/thl/test_survey.py
index 58c4577..e114b70 100644
--- a/tests/managers/thl/test_survey.py
+++ b/tests/managers/thl/test_survey.py
@@ -1,29 +1,41 @@
+from __future__ import annotations
+
import uuid
-from datetime import datetime, timezone
+from collections.abc import Callable
+from datetime import UTC, datetime
from decimal import Decimal
+from typing import TYPE_CHECKING
import pytest
-from generalresearch.models import Source
+from generalresearch.models.definitions 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,
)
+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():
+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=[],
@@ -161,9 +181,10 @@ class TestSurvey:
calc_answers={"i:adhoc_13126": ("3", "4")},
),
]
- uqad = dict()
+ uqad = {}
for uqa in uqas:
- for k, v in uqa.calc_answers.items():
+ assert uqa.calc_answers
+ for k in uqa.calc_answers:
if k in qualifying_questions:
uqad[k] = uqa
uqad[uqa.property_code] = uqa
@@ -174,9 +195,9 @@ class TestSurvey:
qs = sorted(qs, key=lambda x: x.importance.task_count if x.importance else 0)
# qd = {q.id: q for q in qs}
- q = [x for x in qs if x.ext_question_id == "i:adhoc_13126"][0]
+ q = next(x for x in qs if x.ext_question_id == "i:adhoc_13126")
q.explanation_template = "You have been diagnosed with: {answer}."
- q = [x for x in qs if x.ext_question_id == "gr:gender"][0]
+ q = next(x for x in qs if x.ext_question_id == "gr:gender")
q.explanation_template = "Your gender is {answer}."
ecs = []
@@ -205,10 +226,9 @@ class TestSurvey:
class TestSurveyStat:
def test(
self,
- delete_buyers_surveys,
surveystat_manager,
- survey_manager,
- surveys_fixture,
+ survey_manager: SurveyManager,
+ surveys_fixture: list[Survey],
):
survey_manager.create_or_update(surveys_fixture)
ss = [ssa, ssb]
@@ -234,7 +254,7 @@ class TestSurveyStat:
):
survey = surveys_fixture[0].model_copy()
surveys = []
- for idx in range(20_000):
+ for _ in range(20_000):
s = survey.model_copy()
s.survey_id = uuid.uuid4().hex
surveys.append(s)
@@ -251,14 +271,14 @@ class TestSurveyStat:
survey_stats.append(ss)
print(len(survey_stats))
print(survey_stats[12].natural_key, survey_stats[2000].natural_key)
- print(f"----a-----: {datetime.now().isoformat()}")
+ print(f"----a-----: {datetime.now(tz=UTC).isoformat()}")
res = surveystat_manager.update_or_create(survey_stats)
- print(f"----b-----: {datetime.now().isoformat()}")
+ print(f"----b-----: {datetime.now(tz=UTC).isoformat()}")
assert len(res) == 20_000
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
@@ -269,22 +289,24 @@ class TestSurveyStat:
s.conv_beta = 20
s.updated_at = now
# and 1,000 don't change
- print(f"----c-----: {datetime.now().isoformat()}")
+ print(f"----c-----: {datetime.now(tz=UTC).isoformat()}")
res2 = surveystat_manager.update_or_create(ss)
- print(f"----d-----: {datetime.now().isoformat()}")
+ print(f"----d-----: {datetime.now(tz=UTC).isoformat()}")
assert len(res2) == 20_000
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)
@@ -298,14 +320,14 @@ 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
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 4c7dc08..04f69d2 100644
--- a/tests/managers/thl/test_survey_penalty.py
+++ b/tests/managers/thl/test_survey_penalty.py
@@ -1,14 +1,19 @@
+from __future__ import annotations
+
import uuid
+from typing import TYPE_CHECKING
import pytest
-from cachetools.keys import _HashedTuple
-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:
@@ -23,7 +28,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
@@ -49,7 +56,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(
@@ -89,10 +102,8 @@ class TestSurveyPenalty:
)
assert res == {"t:a": 0.1, "t:b": 0.2, "u:b": 0.1}
assert surveypenalty_manager.cache.currsize == 1
- cached_key = tuple(list(list(surveypenalty_manager.cache.keys())[0])[1:])
- assert cached_key == tuple(
- ["product_id", product_uuid, "team_id", team_id_random]
- )
+ cached_key = tuple(list(next(iter(surveypenalty_manager.cache.keys())))[1:])
+ assert cached_key == ("product_id", product_uuid, "team_id", team_id_random)
# Both don't exist, return nothing
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 839bbe1..1e741fd 100644
--- a/tests/managers/thl/test_task_adjustment.py
+++ b/tests/managers/thl/test_task_adjustment.py
@@ -1,27 +1,43 @@
+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
+from typing import TYPE_CHECKING
import pytest
-from datetime import datetime, timezone, timedelta
-from decimal import Decimal
-from generalresearch.models import Source
+from generalresearch.models.definitions import Source
from generalresearch.models.thl.definitions import (
Status,
StatusCode1,
WallAdjustedStatus,
)
+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, Wall
+ from generalresearch.models.thl.user import User
+
@pytest.fixture()
-def session_complete(session_with_tx_factory, 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")
)
@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: User
+):
return session_with_tx_factory(
user=user_with_wallet,
final_status=Status.COMPLETE,
@@ -30,17 +46,21 @@ 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
- )
- wall1 = wall_manager.create_dummy(
- session_id=session.id,
- user_id=user.user_id,
+def session_fail(
+ user: User,
+ session_manager: SessionManager,
+ wall_manager: WallManager,
+ session_factory: Callable[..., Session],
+ wall_factory: Callable[..., Wall],
+) -> Session:
+ session = session_manager.create(started=datetime.now(UTC), user=user)
+ wall1 = wall_factory(
+ session=session,
+ user=user,
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,
@@ -48,46 +68,47 @@ def session_fail(user, session_manager, wall_manager):
status_code_1=StatusCode1.PS_FAIL,
finished=wall1.started + timedelta(seconds=randint(a=60 * 2, b=60 * 10)),
)
- session.wall_events.append(wall1)
return session
class TestHandleRecons:
+ @pytest.fixture(autouse=True)
+ def setup(self, create_main_accounts):
+ create_main_accounts()
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"
+ 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,21 +116,21 @@ 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=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(
@@ -122,7 +143,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 +156,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"
+ 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
)
- amount = ledger_manager.get_account_filtered_balance(
+ commission_account = thl_ledger_manager.get_account_or_create_bp_commission(
+ session_complete_with_wallet.user.product
+ )
+ 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
- ), "earned commission"
+ assert 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
- ), "earned commission"
+ assert 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 55c89c0..11edc99 100644
--- a/tests/managers/thl/test_task_status.py
+++ b/tests/managers/thl/test_task_status.py
@@ -1,47 +1,61 @@
-import pytest
-from datetime import datetime, timezone, timedelta
+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.models import Source
+from generalresearch.models.definitions 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.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=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)
-@pytest.fixture(scope="session")
-def bp1(product_manager):
+@pytest.fixture()
+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):
+@pytest.fixture()
+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(
@@ -52,10 +66,12 @@ def bp2(product_manager):
)
-@pytest.fixture(scope="session")
-def bp3(product_manager):
+@pytest.fixture()
+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(
@@ -70,9 +86,9 @@ class TestTaskStatus:
def test_task_status_complete_1(
self,
- bp1,
- user_factory,
- finished_session_factory,
+ bp1: Product,
+ user_factory: Callable[..., User],
+ finished_session_factory: Callable[..., Session],
session_manager: SessionManager,
):
# User Payout xform NULL
@@ -130,7 +146,11 @@ class TestTaskStatus:
assert tsr == expected_tsr
def test_task_status_complete_2(
- self, bp2, user_factory, finished_session_factory, session_manager
+ self,
+ bp2: Product,
+ user_factory: Callable[..., User],
+ finished_session_factory: Callable[..., Session],
+ session_manager: SessionManager,
):
# User Payout xform 40%
user2: User = user_factory(product=bp2)
@@ -197,7 +217,11 @@ class TestTaskStatus:
assert tsr == expected_tsr
def test_task_status_complete_3(
- self, bp3, user_factory, finished_session_factory, session_manager
+ self,
+ bp3: Product,
+ user_factory: Callable[..., User],
+ 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)
@@ -227,12 +251,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, user_factory, finished_session_factory, session_manager
+ self,
+ bp1: Product,
+ user_factory: Callable[..., User],
+ finished_session_factory: Callable[..., Session],
+ session_manager: SessionManager,
):
# User Payout xform NULL: user payout is None always
user1: User = user_factory(product=bp1)
@@ -263,12 +292,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, user_factory, finished_session_factory, session_manager
+ self,
+ bp2: Product,
+ user_factory: Callable[..., User],
+ finished_session_factory: Callable[..., Session],
+ session_manager: SessionManager,
):
# User Payout xform 40%: user_payout is 0 (not None)
@@ -298,12 +332,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, user_factory, session_factory, session_manager
+ self,
+ bp1: Product,
+ user_factory: Callable[..., User],
+ session_factory: Callable[..., Session],
+ session_manager: SessionManager,
):
# User Payout xform NULL: all payout fields are None
user: User = user_factory(product=bp1)
@@ -332,12 +371,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, user_factory, session_factory, session_manager
+ self,
+ bp2: Product,
+ user_factory: Callable[..., User],
+ 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)
@@ -369,17 +413,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,
- user_factory,
- finished_session_factory,
- wall_manager,
- session_manager,
+ bp1: Product,
+ user_factory: Callable[..., User],
+ 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
@@ -418,17 +463,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,
- user_factory,
- finished_session_factory,
- wall_manager,
- session_manager,
+ bp2: Product,
+ user_factory: Callable[..., User],
+ finished_session_factory: Callable[..., Session],
+ wall_manager: WallManager,
+ session_manager: SessionManager,
):
# Complete -> Fail
# User Payout xform 40%: adjusted_user_payout is 0 (not null)
@@ -470,17 +516,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,
- user_factory,
- session_factory,
- wall_manager,
- session_manager,
+ bp1: Product,
+ user_factory: Callable[..., User],
+ session_factory: Callable[..., Session],
+ wall_manager: WallManager,
+ session_manager: SessionManager,
):
# User Payout xform NULL
user: User = user_factory(product=bp1)
@@ -524,17 +571,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,
- user_factory,
- session_factory,
- wall_manager,
- session_manager,
+ bp2: Product,
+ user_factory: Callable[..., User],
+ session_factory: Callable[..., Session],
+ wall_manager: WallManager,
+ session_manager: SessionManager,
):
# User Payout xform 40%
user: User = user_factory(product=bp2)
@@ -581,17 +629,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,
- user_factory,
- finished_session_factory,
- wall_manager,
- session_manager,
+ bp1: Product,
+ user_factory: Callable[..., User],
+ finished_session_factory: Callable[..., Session],
+ wall_manager: WallManager,
+ session_manager: SessionManager,
):
# User Payout xform NULL
user: User = user_factory(product=bp1)
@@ -635,17 +684,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,
- user_factory,
- finished_session_factory,
- wall_manager,
- session_manager,
+ bp2: Product,
+ user_factory: Callable[..., User],
+ finished_session_factory: Callable[..., Session],
+ wall_manager: WallManager,
+ session_manager: SessionManager,
):
# User Payout xform 40%
user: User = user_factory(product=bp2)
@@ -691,6 +741,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 0d7ffef..5d12052 100644
--- a/tests/managers/thl/test_user_manager/test_base.py
+++ b/tests/managers/thl/test_user_manager/test_base.py
@@ -1,23 +1,37 @@
import logging
-from datetime import datetime, timezone
+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 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,
+)
from generalresearch.managers.thl.user_manager.rate_limit import (
RateLimitItemPerHourConstantKey,
+ UserManagerLimiter,
)
-from generalresearch.managers.thl.user_manager.user_manager import (
- UserManager,
-)
-from generalresearch.models.thl.product import Product, UserCreateConfig
+from generalresearch.models.thl.product import UserCreateConfig
from generalresearch.models.thl.user import User
+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()
@@ -83,10 +97,11 @@ class TestUserManager:
class TestBlockUserManager:
- def test_block_user(self, 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
)
@@ -109,16 +124,19 @@ class TestBlockUserManager:
user = user_manager.get_user(user_id=user.user_id)
assert user.blocked
- def test_block_user_whitelist(self, product, user_manager, thl_web_rw):
+ def test_block_user_whitelist(
+ 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
)
assert not user.blocked
- now = datetime.now(tz=timezone.utc)
+ now = datetime.now(tz=UTC)
# Adds user to whitelist
thl_web_rw.execute_write(
"""
@@ -135,8 +153,13 @@ class TestBlockUserManager:
class TestCreateUserManager:
- def test_create_user(self, product_manager, thl_web_rw, user_manager):
- product: Product = product_manager.create_dummy(
+ def test_create_user(
+ self,
+ product_factory: Callable[..., Product],
+ thl_web_rw: PostgresConfig,
+ user_manager: UserManager,
+ ):
+ product: Product = product_factory(
user_create_config=UserCreateConfig(
min_hourly_create_limit=10, max_hourly_create_limit=69
),
@@ -144,6 +167,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
)
@@ -156,7 +180,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
@@ -172,8 +196,13 @@ class TestCreateUserManager:
assert u2.user_id == user.user_id
assert u2.uuid == user.uuid
- def test_create_user_integrity_error(self, product_manager, user_manager, caplog):
- product: Product = product_manager.create_dummy(
+ def test_create_user_integrity_error(
+ self,
+ user_manager: UserManager,
+ product_factory: Callable[..., Product],
+ caplog,
+ ):
+ product: Product = product_factory(
product_id=uuid4().hex,
team_id=uuid4().hex,
name=f"Test Product ID #{uuid4().hex[:6]}",
@@ -185,6 +214,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(
@@ -213,9 +243,14 @@ class TestCreateUserManager:
assert user1 == user2
- def test_raise_allow_user_create(self, product_manager, user_manager):
+ def test_raise_allow_user_create(
+ 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]}",
@@ -247,10 +282,11 @@ 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
- with pytest.raises(expected_exception=UserCreateNotAllowedError) as cm:
+ with pytest.raises(expected_exception=UserCreateNotAllowedError):
for n, _ in enumerate(range(rl_value + 5)):
user_manager.user_manager_limiter.raise_allow_user_create(
product=product
@@ -260,14 +296,16 @@ class TestCreateUserManager:
class TestUserManagerMethods:
- def test_audit_log(self, user_manager, user, audit_log_manager):
+ def test_audit_log(
+ self, user_manager: UserManager, user: User, audit_log_manager: AuditLogManager
+ ):
from generalresearch.models.thl.userhealth import AuditLog
res = audit_log_manager.filter_by_user_id(user_id=user.user_id)
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_mysql.py b/tests/managers/thl/test_user_manager/test_mysql.py
index 0313bbf..ed7d458 100644
--- a/tests/managers/thl/test_user_manager/test_mysql.py
+++ b/tests/managers/thl/test_user_manager/test_mysql.py
@@ -1,25 +1,28 @@
-from test_utils.models.conftest import user, user_manager
+from __future__ import annotations
+
+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:
- 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 a69519e..f6b59c9 100644
--- a/tests/managers/thl/test_user_manager/test_redis.py
+++ b/tests/managers/thl/test_user_manager/test_redis.py
@@ -1,29 +1,41 @@
+from __future__ import annotations
+
+from typing import TYPE_CHECKING
+
import pytest
from generalresearch.managers.base import Permission
+from generalresearch.managers.thl.user_manager.redis_user_manager import (
+ RedisUserManager,
+)
+
+if TYPE_CHECKING:
+ from generalresearch.config import GRLBaseSettings
+ 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(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 +46,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
@@ -69,9 +87,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 a4b3d57..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,14 +1,25 @@
+from __future__ import annotations
+
+from collections.abc import Callable
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
-from generalresearch.models.thl.user import User
-from test_utils.models.conftest import product, user_manager, user_factory
+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:
- def test_fetch(self, user_factory, product, user_manager):
+ def test_fetch(
+ self,
+ user_factory: Callable[..., User],
+ product: Product,
+ user_manager: UserManager,
+ ):
user1: User = user_factory(product=product)
user2: User = user_factory(product=product)
res = user_manager.fetch_by_bpuids(
@@ -30,7 +41,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 91dc16a..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,20 +1,38 @@
+from __future__ import annotations
+
+from collections.abc import Callable
+from typing import TYPE_CHECKING
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
+
+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:
- 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, product, user_metadata_manager):
- from generalresearch.models.thl.user import User
+ def test_create(
+ self,
+ user_factory: Callable[..., User],
+ product: Product,
+ user_metadata_manager: UserMetadataManager,
+ ):
u1: User = user_factory(product=product)
@@ -27,8 +45,12 @@ 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):
- from generalresearch.models.thl.user import User
+ def test_create_no_email(
+ self,
+ product: Product,
+ user_factory: Callable[..., User],
+ user_metadata_manager: UserMetadataManager,
+ ):
u1: User = user_factory(product=product)
um = UserMetadata(user_id=u1.user_id)
@@ -38,8 +60,12 @@ 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):
- from generalresearch.models.thl.user import User
+ def test_update(
+ self,
+ product: Product,
+ user_factory: Callable[..., User],
+ user_metadata_manager: UserMetadataManager,
+ ):
u: User = user_factory(product=product)
@@ -58,8 +84,9 @@ class TestUserMetadataManager:
email_address=email_address.replace("example1", "example2"),
)
- def test_filter(self, user_factory, product, user_metadata_manager):
- from generalresearch.models.thl.user import User
+ def test_filter(
+ self, user_factory: Callable[..., User], product: Product, user_metadata_manager
+ ):
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 7728f9f..d99b2b8 100644
--- a/tests/managers/thl/test_user_streak.py
+++ b/tests/managers/thl/test_user_streak.py
@@ -1,19 +1,33 @@
+from __future__ import annotations
+
import copy
-from datetime import datetime, timezone, timedelta, date
+from collections.abc import Callable
+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.user_streak import compute_streaks_from_days
-from generalresearch.models.thl.definitions import StatusCode1, Status
+from generalresearch.managers.thl.user_streak import (
+ compute_streaks_from_days,
+)
+from generalresearch.models.thl.definitions import Status, StatusCode1
from generalresearch.models.thl.user_streak import (
- UserStreak,
- StreakState,
- StreakPeriod,
StreakFulfillment,
+ StreakPeriod,
+ StreakState,
+ UserStreak,
)
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.session import SessionManager
+ from generalresearch.managers.thl.user_streak import (
+ UserStreakManager,
+ )
+ from generalresearch.models.thl.session import Session, Wall
+ from generalresearch.models.thl.user import User
+
def test_compute_streaks_from_days():
days = [
@@ -59,7 +73,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,8 +108,12 @@ def broken_active_streak(user):
]
-def create_session_fail(session_manager, start, 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 = session_manager.create(started=start, country_iso="us", user=user)
session_manager.finish_with_status(
session,
finished=start + timedelta(minutes=1),
@@ -104,8 +122,12 @@ def create_session_fail(session_manager, start, user):
)
-def create_session_complete(session_manager, start, user):
- session = session_manager.create_dummy(started=start, country_iso="us", user=user)
+def create_session_complete(
+ session_manager: SessionManager,
+ start: datetime,
+ user: User,
+):
+ session = session_manager.create(started=start, country_iso="us", user=user)
session_manager.finish_with_status(
session,
finished=start + timedelta(minutes=1),
@@ -115,7 +137,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,14 +145,19 @@ 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],
+ bare_session_factory: Callable[..., Session],
+ wall_factory: Callable[..., Wall],
):
# 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
- session = session_manager.create_dummy(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 == []
@@ -171,12 +198,14 @@ 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
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(
@@ -192,16 +221,16 @@ 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
# 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,
@@ -217,9 +246,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 1cda8de..268b110 100644
--- a/tests/managers/thl/test_userhealth.py
+++ b/tests/managers/thl/test_userhealth.py
@@ -1,27 +1,43 @@
-from datetime import timezone, datetime
+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
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,
+)
from generalresearch.models.thl.user_iphistory import (
IPRecord,
+ UserIPHistory,
)
-from generalresearch.models.thl.userhealth import AuditLogLevel, AuditLog
+from generalresearch.models.thl.userhealth import AuditLog, AuditLogLevel
+
+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()
class TestAuditLog:
- def test_init(self, thl_web_rr, 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 +49,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)
@@ -51,14 +68,14 @@ 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,
- user_factory,
- product_factory,
- audit_log_factory,
- audit_log_manager,
+ user_factory: Callable[..., User],
+ product_factory: Callable[..., Product],
+ audit_log_factory: Callable[..., AuditLog],
+ audit_log_manager: AuditLogManager,
):
p1 = product_factory()
p2 = product_factory()
@@ -82,7 +99,11 @@ 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: Callable[..., AuditLog],
+ audit_log_manager: AuditLogManager,
):
u1 = user_factory(product=product)
u2 = user_factory(product=product)
@@ -108,10 +129,10 @@ class TestAuditLog:
def test_filter(
self,
- user_factory,
- product_factory,
- audit_log_factory,
- audit_log_manager,
+ user_factory: Callable[..., User],
+ product_factory: Callable[..., Product],
+ audit_log_factory: Callable[..., AuditLog],
+ audit_log_manager: AuditLogManager,
):
p1 = product_factory()
p2 = product_factory()
@@ -142,10 +163,10 @@ class TestAuditLog:
def test_filter_count(
self,
- user_factory,
- product_factory,
- audit_log_factory,
- audit_log_manager,
+ user_factory: Callable[..., User],
+ product_factory: Callable[..., Product],
+ audit_log_factory: Callable[..., AuditLog],
+ audit_log_manager: AuditLogManager,
):
p1 = product_factory()
p2 = product_factory()
@@ -179,7 +200,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
@@ -205,18 +226,28 @@ class TestAuditLog:
class TestIPRecordManager:
- def test_init(self, thl_web_rr, thl_redis_config, ip_record_manager):
+ 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):
- instance = ip_record_manager.create_dummy(
- user_id=user.user_id, ip=ip_information.ip
- )
+ def test_create(
+ self,
+ ip_record_manager: IPRecordManager,
+ user: User,
+ ip_information: IPInformation,
+ ip_record_factory: Callable[..., IPRecord],
+ ):
+ instance = ip_record_factory(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,20 +259,22 @@ class TestIPRecordManager:
def test_prefetch_info(
self,
- ip_record_factory,
- ip_information_factory,
- ip_geoname,
- user,
- thl_web_rr,
- thl_redis_config,
+ ip_record_factory: Callable[..., IPRecord],
+ ip_information_factory: Callable[..., IPInformation],
+ ip_geoname: IPGeoname,
+ user: User,
+ thl_web_rr: PostgresConfig,
+ thl_redis_config: RedisConfig,
):
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
@@ -265,7 +298,12 @@ 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: UserIpHistoryManager,
+ ):
instance = UserIpHistoryManager(
pg_config=thl_web_rr, redis_config=thl_redis_config
)
@@ -274,27 +312,31 @@ class TestUserIpHistoryManager:
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)
+ 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)
+ 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()
- 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)
+ 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 +345,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 +354,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,23 +369,27 @@ 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
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)
assert len(iph.ips) == 1
ipr = iph.ips[0]
assert ipr.information is not None
@@ -344,23 +397,27 @@ 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
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)
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 ee44e23..777252f 100644
--- a/tests/managers/thl/test_wall_manager.py
+++ b/tests/managers/thl/test_wall_manager.py
@@ -1,25 +1,39 @@
-from datetime import datetime, timezone, timedelta
+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.models import Source
-from generalresearch.models.thl.session import (
+from generalresearch.models.definitions import Source
+from generalresearch.models.thl.definitions import (
ReportValue,
Status,
StatusCode1,
)
-from test_utils.models.conftest import user, session
+
+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, Wall
+ 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
@@ -63,12 +77,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)
@@ -78,21 +95,21 @@ class TestWallManager:
assert isinstance(res, list)
assert len(res) == 50
- res1 = list(set([w.session_id for w in res]))
+ res1 = list({w.session_id for w in res})
res1.sort()
assert session_ids == res1
- def test_create_wall(self, wall_manager, session_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,
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",
- req_cpi=Decimal("1"),
+ req_cpi=Decimal(1),
)
assert w is not None
@@ -100,7 +117,11 @@ class TestWallManager:
assert w == w2
def test_report_wall_abandon(
- self, wall_manager, session_manager, user, session, utc_hour_ago
+ self,
+ wall_manager: WallManager,
+ user: User,
+ session: Session,
+ utc_hour_ago: datetime,
):
w1 = wall_manager.create(
session_id=session.id,
@@ -110,7 +131,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,
@@ -141,7 +162,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,
@@ -151,7 +177,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)
@@ -178,11 +204,15 @@ class TestWallManager:
assert "This survey blows!" == w2.report_notes
def test_filter_wall_attempts(
- self, wall_manager, session_manager, user, session, utc_hour_ago
+ 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
- w1 = wall_manager.create(
+ wall_manager.create(
session_id=session.id,
user_id=user.user_id,
uuid_id=uuid4().hex,
@@ -190,11 +220,11 @@ 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
- w2 = wall_manager.create(
+ wall_manager.create(
session_id=session.id,
user_id=user.user_id,
uuid_id=uuid4().hex,
@@ -202,7 +232,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
@@ -210,21 +240,25 @@ 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,
+ user: User,
+ bare_session_factory: Callable[..., Session],
+ wall_factory: Callable[..., Wall],
):
- 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(
+ session = bare_session_factory(started=start1, user=user)
+ wall_factory(
session_id=session.id,
- user_id=session.user_id,
+ user=session.user,
started=start1,
req_cpi=Decimal("1.23"),
req_survey_id="11111",
@@ -238,9 +272,9 @@ class TestWallCacheManager:
attempts = wall_cache_manager.get_attempts(user_id=user.user_id)
assert len(attempts) == 1
- wall2 = wall_manager.create_dummy(
+ wall_factory(
session_id=session.id,
- user_id=session.user_id,
+ user=session.user,
started=start2,
req_cpi=Decimal("1.23"),
req_survey_id="22222",
@@ -264,10 +298,10 @@ 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)
- wall3 = wall_manager.create_dummy(
+ session = bare_session_factory(started=start3, user=user)
+ wall_factory(
session_id=session.id,
- user_id=session.user_id,
+ user=session.user,
started=start3,
req_cpi=Decimal("1.23"),
req_survey_id="33333",
@@ -279,5 +313,5 @@ class TestWallCacheManager:
redis_key = wall_cache_manager.get_cache_key_(user_id=user.user_id)
assert wall_cache_manager.redis_client.llen(redis_key) == 5000
- assert len(attempts) == 5000
+ assert len(attempts) == 5_000
assert attempts[0].req_survey_id == "33333"
diff --git a/tests/models/admin/test_report_request.py b/tests/models/admin/test_report_request.py
index a80afbe..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 datetime, timezone
+from datetime import UTC, datetime
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..e8a5aa3 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
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())
@@ -42,7 +42,7 @@ class TestAwareDatetimeISO:
AwareDatetimeISOModel.model_validate_json(t.model_dump_json())
def test_no_tz(self):
- dt = datetime(2023, 10, 10, 1, 1, 1)
+ dt = datetime(2023, 10, 10, 1, 1, 1) # noqa
with pytest.raises(expected_exception=ValidationError):
AwareDatetimeISOModel(dt=dt, dt_optional=None)
diff --git a/tests/models/custom_types/test_dsn.py b/tests/models/custom_types/test_dsn.py
index 16e1f83..eff02d3 100644
--- a/tests/models/custom_types/test_dsn.py
+++ b/tests/models/custom_types/test_dsn.py
@@ -1,4 +1,5 @@
-from typing import Optional
+from __future__ import annotations
+
from uuid import uuid4
import pytest
@@ -11,16 +12,15 @@ 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 ---
class TestDaskDsn:
-
def test_base(self):
from dask.distributed import Client
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 736c971..b3a9f13 100644
--- a/tests/models/dynata/test_eligbility.py
+++ b/tests/models/dynata/test_eligbility.py
@@ -1,4 +1,6 @@
-from datetime import datetime, timezone
+from __future__ import annotations
+
+from datetime import UTC, datetime
class TestEligibility:
@@ -40,7 +42,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 +174,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/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 6c84a5d..21e07a4 100644
--- a/tests/models/gr/test_authentication.py
+++ b/tests/models/gr/test_authentication.py
@@ -1,21 +1,30 @@
+from __future__ import annotations
+
import binascii
import json
import os
-from datetime import datetime, timezone
+from collections.abc import Callable
+from datetime import UTC, datetime
from random import randint
-from typing import Callable
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
-from generalresearch.models.gr.authentication import GRUser
-from generalresearch.models.gr.team import Membership, Team
+from generalresearch.models.gr.authentication import Claims, GRToken, GRUser
+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 = ""
class TestGRUser:
-
def test_init(self, gr_user: GRUser):
assert isinstance(gr_user, GRUser)
@@ -29,7 +38,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,
+ gr_membership: Membership,
+ gr_db: PostgresConfig,
+ gr_redis_config: RedisConfig,
+ ):
assert gr_user.teams is None
@@ -41,18 +56,18 @@ class TestGRUser:
def test_prefetch_team_duplicates(
self,
- gr_user_token,
+ gr_user_token: GRToken,
gr_user: GRUser,
- membership: Membership,
- product_factory,
- membership_factory,
- team: Team,
- thl_web_rr,
- gr_redis_config,
- gr_db,
+ gr_membership: Membership,
+ product_factory: Callable[..., Product],
+ 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(gr_team=gr_team, gr_user=gr_user)
gr_user.prefetch_teams(
pg_config=gr_db,
@@ -64,12 +79,12 @@ class TestGRUser:
def test_products(
self,
gr_user: GRUser,
- product_factory,
- team: Team,
- membership: Membership,
- gr_db,
- thl_web_rr,
- gr_redis_config,
+ product_factory: Callable[..., Product],
+ gr_team: Team,
+ gr_membership: Membership,
+ gr_db: PostgresConfig,
+ thl_web_rr: PostgresConfig,
+ gr_redis_config: RedisConfig,
):
from generalresearch.models.thl.product import Product
@@ -77,13 +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)
- p: Product = product_factory(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=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,
@@ -96,8 +113,7 @@ class TestGRUser:
class TestGRUserMethods:
-
- def test_cache_key(self, gr_user, gr_redis):
+ 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
@@ -105,14 +121,13 @@ class TestGRUserMethods:
def test_to_redis(
self,
gr_user: GRUser,
- gr_redis,
- team: Team,
- business,
- product_factory,
- membership_factory: Callable[Membership],
+ gr_team: Team,
+ gr_business: Business,
+ product_factory: Callable[..., Product],
+ gr_membership_factory: Callable[..., Membership],
):
- product_factory(team=team, business=business)
- membership_factory(team=team, gr_user=gr_user)
+ product_factory(team=gr_team, business=gr_business)
+ gr_membership_factory(gr_team=gr_team, gr_user=gr_user)
res = gr_user.to_redis()
assert isinstance(res, str)
@@ -125,49 +140,50 @@ class TestGRUserMethods:
def test_set_cache(
self,
gr_user: GRUser,
- gr_user_token,
- gr_redis,
- gr_db,
- thl_web_rr,
- gr_redis_config,
+ 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,
- gr_redis,
- gr_redis_config,
- gr_db,
- thl_web_rr,
- product_factory,
- team,
- membership_factory,
- thl_redis_config,
+ gr_redis_config: RedisConfig,
+ gr_db: PostgresConfig,
+ thl_web_rr: PostgresConfig,
+ product_factory: Callable[..., Product],
+ gr_team: Team,
+ gr_membership_factory: Callable[..., Membership],
+ thl_redis_config: RedisConfig,
):
from generalresearch.models.gr.authentication import GRUser
- p1 = product_factory(team=team)
- membership_factory(team=team, gr_user=gr_user)
+ client = gr_redis_config.create_redis_client()
+
+ p1 = product_factory(team=gr_team)
+ 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
)
- 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(
@@ -183,22 +199,21 @@ class TestGRUserMethods:
def test_set_cache_team_uuids(
self,
- gr_user,
- membership,
- gr_user_token,
- gr_redis,
- gr_db,
- thl_web_rr,
- product_factory,
- team,
- gr_redis_config,
+ gr_user: GRUser,
+ gr_db: PostgresConfig,
+ thl_web_rr: PostgresConfig,
+ product_factory: Callable[..., Product],
+ gr_team: Team,
+ gr_redis_config: RedisConfig,
+ gr_membership,
):
- product_factory(team=team)
+ product_factory(team=gr_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
@@ -206,81 +221,74 @@ class TestGRUserMethods:
def test_set_cache_business_uuids(
self,
gr_user: GRUser,
- gr_redis,
- gr_db,
- thl_web_rr,
- product_factory,
- business,
- team,
- gr_redis_config,
+ gr_db: PostgresConfig,
+ thl_web_rr: PostgresConfig,
+ product_factory: Callable[..., Product],
+ gr_business: Business,
+ gr_team: Team,
+ gr_redis_config: RedisConfig,
):
- product_factory(team=team, business=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
)
- 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,
- membership,
- gr_user_token,
- gr_redis,
- gr_db,
- thl_web_rr,
- product_factory,
- team,
- gr_redis_config,
+ gr_user: GRUser,
+ gr_db: PostgresConfig,
+ thl_web_rr: PostgresConfig,
+ product_factory: Callable[..., Product],
+ gr_team: Team,
+ gr_redis_config: RedisConfig,
+ gr_membership,
):
- 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
)
- 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
class TestGRToken:
-
@pytest.fixture
- def gr_token(self, gr_user):
- from generalresearch.models.gr.authentication import GRToken
-
- now = datetime.now(tz=timezone.utc)
+ def gr_token(self, gr_user: GRUser):
+ now = datetime.now(tz=UTC)
token = binascii.hexlify(os.urandom(20)).decode()
gr_token = GRToken(key=token, created=now, user_id=gr_user.id)
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 a9f01a8..5ab5dff 100644
--- a/tests/models/gr/test_base.py
+++ b/tests/models/gr/test_base.py
@@ -1,16 +1,20 @@
+from __future__ import annotations
+
import subprocess
+from collections.abc import Callable
from pathlib import Path
-from typing import Callable
+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:
- def test_git(self, git_key_path: Path, gr_repo: Callable[..., Path]):
+ def test_git(self, gr_repo: Callable[..., Path]):
repo_path = gr_repo()
try:
@@ -33,14 +37,5 @@ 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, 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
diff --git a/tests/models/gr/test_business.py b/tests/models/gr/test_business.py
index 7a84f23..e38850d 100644
--- a/tests/models/gr/test_business.py
+++ b/tests/models/gr/test_business.py
@@ -1,7 +1,11 @@
+from __future__ import annotations
+
import os
-from datetime import datetime, timedelta, timezone
+from collections.abc import Callable
+from datetime import UTC, datetime, timedelta
from decimal import Decimal
-from typing import Optional
+from pathlib import Path
+from typing import TYPE_CHECKING
from uuid import uuid4
import pandas as pd
@@ -15,34 +19,55 @@ from distributed.utils_test import (
from pytest import approx
from generalresearch.currency import USDCent
-from generalresearch.managers.gr.business import BusinessBankAccountManager
from generalresearch.models.gr.business import (
Business,
BusinessAddress,
- BusinessBankAccount,
BusinessContact,
)
from generalresearch.models.thl.finance import (
BusinessBalances,
ProductBalances,
)
-from generalresearch.pg_helper import PostgresConfig
+from generalresearch.models.thl.product import Product
+
+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:
-
def test_init(
self,
- business: Business,
- business_bank_account_manager: BusinessBankAccountManager,
+ gr_business: Business,
+ gr_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=business.id,
+ instance = gr_business_bank_account_manager.create(
+ business_id=gr_business.id,
uuid=uuid4().hex,
transfer_method=TransferMethod.ACH,
)
@@ -50,30 +75,28 @@ class TestBusinessBankAccount:
def test_business(
self,
- business_bank_account: BusinessBankAccount,
- business: Business,
- gr_db,
- gr_redis_config,
+ 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 == business.uuid
+ assert isinstance(gr_business_bank_account.business, Business)
+ assert gr_business_bank_account.business.uuid == gr_business.uuid
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")
@@ -82,346 +105,352 @@ 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"
+ return "30D"
@pytest.fixture
- def duration(self) -> Optional["timedelta"]:
+ def duration(self) -> timedelta | None:
return None
- def test_init(self, business):
- from generalresearch.models.gr.business import 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,
- business,
- product_factory,
- thl_web_rr,
- lm,
- thl_lm,
- business_payout_event_manager,
- bp_payout_factory,
- start,
- user_factory,
- session_with_tx_factory,
- pop_ledger_merge,
+ 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,
+ brokerage_product_payout_event_factory: Callable[
+ ..., BusinessPayoutEventManager
+ ],
+ start: datetime,
+ user_factory: Callable[..., User],
+ session_with_tx_factory: Callable[..., Session],
+ pop_ledger_merge: PopLedgerMerge,
client_no_amm: DaskClient,
ledger_collection,
- mnt_filepath,
- create_main_accounts,
+ mnt_filepath: GRLDatasets,
+ 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)
- thl_lm.get_account_or_create_bp_wallet(product=p1)
- thl_lm.get_account_or_create_bp_wallet(product=p2)
+ 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 "<Business: " in res1
- res2 = str(business)
+ res2 = str(gr_business)
- assert business.uuid in res2
+ assert gr_business.uuid in res2
assert "Name:" in res2
assert "Not Loaded" in res2
- business.prefetch_products(thl_pg_config=thl_web_rr)
- business.prefetch_bp_accounts(thl_lm=thl_lm, thl_pg_config=thl_web_rr)
- res3 = str(business)
+ gr_business.prefetch_products(product_manager=product_manager)
+ gr_business.prefetch_bp_accounts(
+ thl_lm=thl_ledger_manager, product_manager=product_manager
+ )
+ res3 = str(gr_business)
assert "Products: 2" in res3
assert "Ledger Accounts: 2" in res3
# -- need some tx to make these interesting
- business_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
session_with_tx_factory(
user=u1,
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),
- 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)
- business.prebuild_payouts(
- thl_pg_config=thl_web_rr,
- thl_lm=thl_lm,
+ gr_business.prebuild_payouts(
bpem=business_payout_event_manager,
)
- business.prebuild_balance(
- thl_pg_config=thl_web_rr,
- lm=lm,
+ gr_business.prebuild_balance(
+ product_manager=product_manager,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
)
- res4 = str(business)
+ res4 = str(gr_business)
assert "Payouts: 1" in res4
assert "Available Balance: 141" in res4
- def test_addresses(self, business, business_address, gr_db):
+ def test_addresses(
+ self, gr_business: Business, gr_db: PostgresConfig, gr_business_address
+ ):
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, business, team, team_manager, gr_db):
- assert business.teams is None
+ def test_teams(
+ self,
+ gr_business: Business,
+ gr_team: Team,
+ gr_team_manager: TeamManager,
+ gr_db: PostgresConfig,
+ ):
+ 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
+ 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
- def test_products(self, business, product_factory, thl_web_rr):
- from generalresearch.models.thl.product import Product
+ def test_products(
+ self,
+ 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(thl_pg_config=thl_web_rr)
- 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
- p2 = product_factory(business=business)
- p3 = 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(thl_pg_config=thl_web_rr)
- assert len(business.products) == 3
+ gr_business.prefetch_products(product_manager=product_manager)
+ assert len(gr_business.products) == 3
- def test_bank_accounts(self, business, business_bank_account, gr_db):
- assert business.products is None
+ def test_bank_accounts(
+ self,
+ gr_business: Business,
+ gr_business_bank_account,
+ gr_business_bank_account_manager: BusinessBankAccountManager,
+ ):
+ assert gr_business.products is None
# It's an empty list after prefetch
- business.prefetch_bank_accounts(pg_config=gr_db)
- assert isinstance(business.bank_accounts, list)
- assert len(business.bank_accounts) == 1
+ gr_business.prefetch_bank_accounts(
+ business_bank_account_manager=gr_business_bank_account_manager
+ )
+ assert isinstance(gr_business.bank_accounts, list)
+ assert len(gr_business.bank_accounts) == 1
def test_balance(
self,
- business: Business,
- mnt_filepath,
+ gr_business: Business,
+ mnt_filepath: GRLDatasets,
client_no_amm: DaskClient,
thl_web_rr: PostgresConfig,
- ledger_manager,
- pop_ledger_merge,
+ ledger_manager: LedgerManager,
+ pop_ledger_merge: PopLedgerMerge,
+ product_manager: ProductManager,
):
- assert business.balance is None
+ assert gr_business.balance is None
with pytest.raises(expected_exception=AssertionError) as cm:
- business.prebuild_balance(
- thl_pg_config=thl_web_rr,
+ gr_business.prebuild_balance(
+ product_manager=product_manager,
lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
)
assert "Cannot build Business Balance" in str(cm.value)
- assert business.balance is None
+ assert gr_business.balance is None
# TODO: Add parquet building so that this doesn't fail and we can
# properly assign a business.balance
def test_payouts_no_accounts(
self,
- business,
- product_factory,
- thl_web_rr,
- thl_ledger_manager,
- business_payout_event_manager,
+ gr_business: Business,
+ product_factory: Callable[..., Product],
+ thl_ledger_manager: ThlLedgerManager,
+ business_payout_event_manager: BusinessPayoutEventManager,
):
- assert business.payouts is None
+ assert gr_business.payouts is None
with pytest.raises(expected_exception=AssertionError) as cm:
- business.prebuild_payouts(
- thl_pg_config=thl_web_rr,
- thl_lm=thl_ledger_manager,
+ gr_business.prebuild_payouts(
bpem=business_payout_event_manager,
)
assert "Must provide product_uuids" in str(cm.value)
- p = product_factory(business=business)
+ p = product_factory(business=gr_business)
thl_ledger_manager.get_account_or_create_bp_wallet(product=p)
- 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) == 0
+ assert isinstance(gr_business.payouts, list)
+ assert len(gr_business.payouts) == 0
def test_payouts(
self,
- business: Business,
- product_factory: Callable[Product],
- bp_payout_factory,
- thl_ledger_manager,
- thl_web_rr,
- business_payout_event_manager,
- create_main_accounts,
+ gr_business: Business,
+ product_factory: Callable[..., Product],
+ brokerage_product_payout_event_factory: Callable[
+ ..., BrokerageProductPayoutEvent
+ ],
+ thl_ledger_manager: ThlLedgerManager,
+ business_payout_event_manager: BusinessPayoutEventManager,
+ create_main_accounts: Callable[..., None],
):
create_main_accounts()
- p = product_factory(business=business)
+ 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_lm)
- bp_payout_factory(
- product=p, amount=USDCent(123), skip_wallet_balance_check=True
- )
+ brokerage_product_payout_event_factory(product=p, amount=USDCent(123))
- business.prebuild_payouts(
- thl_pg_config=thl_web_rr,
- thl_lm=thl_ledger_manager,
+ gr_business.prebuild_payouts(
bpem=business_payout_event_manager,
)
- assert len(business.payouts) == 1
- assert sum([p.amount for p in business.payouts]) == 123
+ assert len(gr_business.payouts) == 1
+ 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,
- skip_one_per_day_check=True,
- )
- business_payout_event_manager.set_account_lookup_table(
- thl_lm=thl_ledger_manager
)
- business.prebuild_payouts(
- thl_pg_config=thl_web_rr,
- thl_lm=thl_ledger_manager,
+ gr_business.prebuild_payouts(
bpem=business_payout_event_manager,
)
- assert len(business.payouts) == 1
- assert len(business.payouts[0].bp_payouts) == 2
- assert sum([p.amount for p in business.payouts]) == 246
+ assert isinstance(gr_business.payouts, list)
+ 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(
self,
- business,
- product_factory,
- bp_payout_factory,
- thl_lm,
- thl_web_rr,
- business_payout_event_manager,
- create_main_accounts,
+ gr_business: Business,
+ product_factory: Callable[..., Product],
+ brokerage_product_payout_event_factory: Callable[
+ ..., BrokerageProductPayoutEvent
+ ],
+ thl_ledger_manager: ThlLedgerManager,
+ thl_web_rr: PostgresConfig,
+ business_payout_event_manager: BusinessPayoutEventManager,
+ create_main_accounts: Callable[..., None],
):
- from generalresearch.models.thl.product import Product
create_main_accounts()
- p1: Product = product_factory(business=business)
- thl_lm.get_account_or_create_bp_wallet(product=p1)
- business_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
+ p1: Product = product_factory(business=gr_business)
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=p1)
- 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,
- skip_one_per_day_check=True,
)
- business.prebuild_payouts(
- thl_pg_config=thl_web_rr,
- thl_lm=thl_lm,
+ gr_business.prebuild_payouts(
bpem=business_payout_event_manager,
)
- assert len(business.payouts) == 1
- assert len(business.payouts[0].bp_payouts) == 3
- assert business.payouts_total == USDCent(76)
- assert business.payouts_total_str == "$0.76"
+ assert isinstance(gr_business.payouts, list)
+ 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"
def test_pop_financial(
self,
- business,
- thl_web_rr,
- thl_ledger_manager,
- mnt_filepath,
- client_no_amm,
- pop_ledger_merge,
+ gr_business: Business,
+ product_manager: ProductManager,
+ thl_ledger_manager: ThlLedgerManager,
+ mnt_filepath: GRLDatasets,
+ client_no_amm: DaskClient,
+ pop_ledger_merge: PopLedgerMerge,
):
- assert business.pop_financial is None
- business.prebuild_pop_financial(
- thl_pg_config=thl_web_rr,
- thl_lm=thl_lm,
+ assert gr_business.pop_financial is None
+ gr_business.prebuild_pop_financial(
+ product_manager=product_manager,
+ thl_lm=thl_ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
)
- assert business.pop_financial == []
+ assert gr_business.pop_financial == []
- def test_bp_accounts(self, business, lm, thl_web_rr, product_factory, thl_lm):
- assert business.bp_accounts is None
- business.prefetch_bp_accounts(thl_lm=thl_lm, thl_pg_config=thl_web_rr)
- assert business.bp_accounts == []
-
- from generalresearch.models.thl.product import Product
+ def test_bp_accounts(
+ self,
+ gr_business: Business,
+ product_factory: Callable[..., Product],
+ thl_ledger_manager: ThlLedgerManager,
+ product_manager: ProductManager,
+ ):
+ assert gr_business.bp_accounts is None
+ gr_business.prefetch_bp_accounts(
+ thl_lm=thl_ledger_manager, product_manager=product_manager
+ )
+ assert gr_business.bp_accounts == []
- p1: Product = product_factory(business=business)
- thl_lm.get_account_or_create_bp_wallet(product=p1)
+ p1: Product = product_factory(business=gr_business)
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=p1)
- business.prefetch_bp_accounts(thl_lm=thl_lm, thl_pg_config=thl_web_rr)
- assert len(business.bp_accounts) == 1
+ gr_business.prefetch_bp_accounts(
+ thl_lm=thl_ledger_manager, product_manager=product_manager
+ )
+ assert len(gr_business.bp_accounts) == 1
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"
+ return "30D"
@pytest.fixture
- def duration(self) -> Optional["timedelta"]:
+ def duration(self) -> timedelta | None:
return None
@pytest.mark.skip
@@ -432,34 +461,27 @@ class TestBusinessBalance:
def test_single_product(
self,
- business,
- product_factory,
- user_factory,
+ gr_business: Business,
+ product_factory: Callable[..., Product],
+ user_factory: Callable[..., User],
mnt_filepath,
- bp_payout_factory,
- thl_lm,
- lm,
- duration,
- offset,
- start,
- thl_web_rr,
- payout_event_manager,
- session_with_tx_factory,
- delete_ledger_db,
- create_main_accounts,
- client_no_amm,
+ ledger_manager: LedgerManager,
+ start: datetime,
+ thl_web_rr: PostgresConfig,
+ session_with_tx_factory: Callable[..., Session],
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ client_no_amm: DaskClient,
ledger_collection,
- pop_ledger_merge,
- delete_df_collection,
+ product_manager: ProductManager,
+ pop_ledger_merge: PopLedgerMerge,
+ delete_df_collection: Callable[..., None],
):
delete_ledger_db()
create_main_accounts()
delete_df_collection(coll=ledger_collection)
- from generalresearch.models.thl.product import Product
- from generalresearch.models.thl.user import User
-
- p1: Product = product_factory(business=business)
+ p1: Product = product_factory(business=gr_business)
u1: User = user_factory(product=p1)
u2: User = user_factory(product=p1)
@@ -478,57 +500,50 @@ 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(
- thl_pg_config=thl_web_rr,
- lm=lm,
+ gr_business.prebuild_balance(
+ product_manager=product_manager,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
)
- assert isinstance(business.balance, BusinessBalances)
- assert business.balance.payout == 190
- assert business.balance.adjustment == 0
- assert business.balance.net == 190
- assert business.balance.retainer == 47
- assert business.balance.available_balance == 143
+ assert isinstance(gr_business.balance, BusinessBalances)
+ assert gr_business.balance.payout == 190
+ assert gr_business.balance.adjustment == 0
+ assert gr_business.balance.net == 190
+ assert gr_business.balance.retainer == 47
+ assert gr_business.balance.available_balance == 143
- assert len(business.balance.product_balances) == 1
- pb = business.balance.product_balances[0]
+ assert len(gr_business.balance.product_balances) == 1
+ pb = gr_business.balance.product_balances[0]
assert isinstance(pb, ProductBalances)
- assert pb.balance == business.balance.balance
- assert pb.available_balance == business.balance.available_balance
+ assert pb.balance == gr_business.balance.balance
+ assert pb.available_balance == gr_business.balance.available_balance
assert pb.adjustment_percent == 0.0
def test_multi_product(
self,
- business,
- product_factory,
- user_factory,
- mnt_filepath,
- bp_payout_factory,
- thl_lm,
- ledger_manager,
- duration,
- offset,
- start,
- thl_web_rr,
- payout_event_manager,
- session_with_tx_factory,
- delete_ledger_db,
- create_main_accounts,
- client_no_amm,
+ gr_business: Business,
+ product_factory: Callable[..., Product],
+ user_factory: Callable[..., User],
+ mnt_filepath: GRLDatasets,
+ ledger_manager: LedgerManager,
+ product_manager: ProductManager,
+ start: datetime,
+ session_with_tx_factory: Callable[..., Session],
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ client_no_amm: DaskClient,
ledger_collection,
- pop_ledger_merge,
- delete_df_collection,
+ pop_ledger_merge: PopLedgerMerge,
+ delete_df_collection: Callable[..., None],
):
delete_ledger_db()
create_main_accounts()
delete_df_collection(coll=ledger_collection)
- from generalresearch.models.thl.user import User
-
- u1: User = user_factory(product=product_factory(business=business))
- u2: 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))
session_with_tx_factory(
user=u1,
@@ -545,33 +560,33 @@ 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(
- thl_pg_config=thl_web_rr,
- lm=lm,
+ gr_business.prebuild_balance(
+ product_manager=product_manager,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
)
- assert isinstance(business.balance, BusinessBalances)
- assert business.balance.payout == 190
- assert business.balance.balance == 190
- assert business.balance.adjustment == 0
- assert business.balance.net == 190
- assert business.balance.retainer == 46
- assert business.balance.available_balance == 144
+ assert isinstance(gr_business.balance, BusinessBalances)
+ assert gr_business.balance.payout == 190
+ assert gr_business.balance.balance == 190
+ assert gr_business.balance.adjustment == 0
+ assert gr_business.balance.net == 190
+ assert gr_business.balance.retainer == 46
+ assert gr_business.balance.available_balance == 144
- assert len(business.balance.product_balances) == 2
+ assert len(gr_business.balance.product_balances) == 2
- pb1 = business.balance.product_balances[0]
- pb2 = business.balance.product_balances[1]
+ pb1 = gr_business.balance.product_balances[0]
+ pb2 = gr_business.balance.product_balances[1]
assert isinstance(pb1, ProductBalances)
assert pb1.product_id == u1.product_id
assert isinstance(pb2, ProductBalances)
assert pb2.product_id == u2.product_id
for pb in [pb1, pb2]:
- assert pb.balance != business.balance.balance
- assert pb.available_balance != business.balance.available_balance
+ assert pb.balance != gr_business.balance.balance
+ assert pb.available_balance != gr_business.balance.available_balance
assert pb.adjustment_percent == 0.0
assert pb1.product_id in [u1.product_id, u2.product_id]
@@ -592,34 +607,33 @@ class TestBusinessBalance:
def test_multi_product_multi_payout(
self,
- business,
- product_factory,
- user_factory,
- mnt_filepath,
- bp_payout_factory,
- thl_lm,
- lm,
- duration,
- offset,
- start,
- thl_web_rr,
- payout_event_manager,
- session_with_tx_factory,
- delete_ledger_db,
- create_main_accounts,
- client_no_amm,
+ gr_business: Business,
+ product_factory: Callable[..., Product],
+ user_factory: Callable[..., User],
+ product_manager: ProductManager,
+ mnt_filepath: GRLDatasets,
+ brokerage_product_payout_event_factory: Callable[
+ ..., BrokerageProductPayoutEvent
+ ],
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ start: datetime,
+ thl_web_rr: PostgresConfig,
+ payout_event_manager: PayoutEventManager,
+ session_with_tx_factory: Callable[..., None],
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ client_no_amm: DaskClient,
ledger_collection,
- pop_ledger_merge,
- delete_df_collection,
+ pop_ledger_merge: PopLedgerMerge,
+ delete_df_collection: Callable[..., None],
):
delete_ledger_db()
create_main_accounts()
delete_df_collection(coll=ledger_collection)
- from generalresearch.models.thl.user import User
-
- u1: User = user_factory(product=product_factory(business=business))
- u2: 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))
session_with_tx_factory(
user=u1,
@@ -633,62 +647,58 @@ class TestBusinessBalance:
started=start + timedelta(days=2),
)
- payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
-
- bp_payout_factory(
+ brokerage_product_payout_event_factory(
product=u1.product,
amount=USDCent(5),
created=start + timedelta(days=4),
- skip_wallet_balance_check=True,
- 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),
- 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)
- business.prebuild_balance(
- thl_pg_config=thl_web_rr,
- lm=lm,
+ gr_business.prebuild_balance(
+ product_manager=product_manager,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
)
- assert business.balance.payout == 190
- assert business.balance.net == 190
+ assert isinstance(gr_business.balance, BusinessBalances)
+ assert gr_business.balance.payout == 190
+ assert gr_business.balance.net == 190
- assert business.balance.balance == 135
+ assert gr_business.balance.balance == 135
def test_multi_product_multi_payout_adjustment(
self,
- business,
- product_factory,
- user_factory,
- mnt_filepath,
- bp_payout_factory,
- duration,
- offset,
- start,
- thl_web_rr,
- payout_event_manager,
- session_with_tx_factory,
- delete_ledger_db,
- create_main_accounts,
+ gr_business: Business,
+ product_factory: Callable[..., Product],
+ user_factory: Callable[..., User],
+ mnt_filepath: GRLDatasets,
+ brokerage_product_payout_event_factory: Callable[
+ ..., BrokerageProductPayoutEvent
+ ],
+ ledger_manager: LedgerManager,
+ thl_ledger_manager: ThlLedgerManager,
+ start: datetime,
+ thl_web_rr: PostgresConfig,
+ 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,
- pop_ledger_merge,
- wall_manager,
- session_manager,
- adj_to_fail_with_tx_factory,
- delete_df_collection,
+ pop_ledger_merge: PopLedgerMerge,
+ adj_to_fail_with_tx_factory: Callable[..., None],
+ delete_df_collection: Callable[..., None],
):
"""
- Product 1 $2.50 Complete
@@ -712,11 +722,9 @@ class TestBusinessBalance:
delete_df_collection(coll=ledger_collection)
delete_df_collection(coll=task_adj_collection)
- from generalresearch.models.thl.user import User
-
- 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,
@@ -729,22 +737,17 @@ class TestBusinessBalance:
wall_req_cpi=Decimal("2.50"),
started=start + timedelta(days=2),
)
- payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
- bp_payout_factory(
+ brokerage_product_payout_event_factory(
product=u1.product,
amount=USDCent(250),
created=start + timedelta(days=3),
- skip_wallet_balance_check=True,
- 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),
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
)
adj_to_fail_with_tx_factory(session=s1, created=start + timedelta(days=5))
@@ -769,57 +772,60 @@ class TestBusinessBalance:
df = client_no_amm.compute(pop_ledger_merge.ddf(), sync=True)
assert df.shape == (20, 28)
- business.prebuild_balance(
- thl_pg_config=thl_web_rr,
- lm=lm,
+ gr_business.prebuild_balance(
+ product_manager=product_manager,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
)
- assert business.balance.payout == 714
- assert business.balance.adjustment == -238
+ assert isinstance(gr_business.balance, BusinessBalances)
+ assert gr_business.balance.payout == 714
+ assert gr_business.balance.adjustment == -238
- assert business.balance.product_balances[0].adjustment == -238
- assert business.balance.product_balances[1].adjustment == 0
- assert business.balance.product_balances[2].adjustment == 0
+ assert gr_business.balance.product_balances[0].adjustment == -238
+ assert gr_business.balance.product_balances[1].adjustment == 0
+ assert gr_business.balance.product_balances[2].adjustment == 0
- assert business.balance.expense == 0
- assert business.balance.net == 714 - 238
- assert business.balance.balance == business.balance.payout - (250 + 50 + 238)
+ assert gr_business.balance.expense == 0
+ assert gr_business.balance.net == 714 - 238
+ assert gr_business.balance.balance == gr_business.balance.payout - (
+ 250 + 50 + 238
+ )
predicted_retainer = sum(
[
pb.balance * 0.25
- for pb in business.balance.product_balances
+ for pb in gr_business.balance.product_balances
if pb.balance > 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,
- product,
- mnt_filepath,
- thl_lm,
- client_no_amm,
- thl_redis_config,
- brokerage_product_payout_event_manager,
- delete_ledger_db,
- create_main_accounts,
- delete_df_collection,
+ mnt_filepath: GRLDatasets,
+ thl_ledger_manager: ThlLedgerManager,
+ client_no_amm: DaskClient,
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ delete_df_collection: Callable[..., None],
ledger_collection,
- business,
- user_factory,
- product_factory,
- session_with_tx_factory,
- pop_ledger_merge,
- start,
- bp_payout_factory,
+ gr_business: Business,
+ user_factory: Callable[..., User],
+ product_factory: Callable[..., Product],
+ session_with_tx_factory: Callable[..., Session],
+ pop_ledger_merge: PopLedgerMerge,
+ start: datetime,
+ brokerage_product_payout_event_factory: Callable[
+ ..., BrokerageProductPayoutEvent
+ ],
payout_event_manager,
- adj_to_fail_with_tx_factory,
- thl_web_rr,
- lm,
+ product_manager: ProductManager,
+ adj_to_fail_with_tx_factory: Callable[..., None],
+ thl_web_rr: PostgresConfig,
+ ledger_manager: LedgerManager,
):
"""Test having a Business with two products.. one that lost money
and one that gained money. Ensure that the Business balance
@@ -830,15 +836,12 @@ class TestBusinessBalance:
create_main_accounts()
delete_df_collection(coll=ledger_collection)
- from generalresearch.models.thl.product import Product
- from generalresearch.models.thl.user import User
-
- 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_lm.get_account_or_create_bp_wallet(product=p1)
- thl_lm.get_account_or_create_bp_wallet(product=p2)
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=p1)
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=p2)
# Product 1: Complete, Payout, Recon..
s1 = session_with_tx_factory(
@@ -846,14 +849,11 @@ class TestBusinessBalance:
wall_req_cpi=Decimal(".75"),
started=start + timedelta(days=1),
)
- payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
- bp_payout_factory(
+ brokerage_product_payout_event_factory(
product=u1.product,
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,
@@ -861,12 +861,12 @@ class TestBusinessBalance:
)
# Product 2: Complete, Complete.
- s2 = session_with_tx_factory(
+ session_with_tx_factory(
user=u2,
wall_req_cpi=Decimal(".75"),
started=start + timedelta(days=1, minutes=3),
)
- s3 = session_with_tx_factory(
+ session_with_tx_factory(
user=u2,
wall_req_cpi=Decimal(".75"),
started=start + timedelta(days=1, minutes=4),
@@ -876,16 +876,17 @@ 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(
- thl_pg_config=thl_web_rr,
- lm=lm,
+ gr_business.prebuild_balance(
+ product_manager=product_manager,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
)
# Check Product 1
- 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
@@ -895,7 +896,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
@@ -905,7 +906,8 @@ 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
assert bb1.expense == 0
@@ -923,29 +925,27 @@ class TestBusinessBalance:
def test_multi_product_multi_payout_adjustment_at_timestamp(
self,
- business,
- product_factory,
- user_factory,
- mnt_filepath,
- bp_payout_factory,
- thl_lm,
- lm,
- duration,
- offset,
- start,
- thl_web_rr,
+ gr_business: Business,
+ product_factory: Callable[..., Product],
+ user_factory: Callable[..., User],
+ mnt_filepath: GRLDatasets,
+ brokerage_product_payout_event_factory: Callable[
+ ..., BrokerageProductPayoutEvent
+ ],
+ thl_ledger_manager: ThlLedgerManager,
+ ledger_manager: LedgerManager,
+ product_manager: ProductManager,
+ start: datetime,
payout_event_manager,
- session_with_tx_factory,
- delete_ledger_db,
- create_main_accounts,
- client_no_amm,
+ session_with_tx_factory: Callable[..., None],
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ client_no_amm: DaskClient,
ledger_collection,
task_adj_collection,
- pop_ledger_merge,
- wall_manager,
- session_manager,
- adj_to_fail_with_tx_factory,
- delete_df_collection,
+ pop_ledger_merge: PopLedgerMerge,
+ adj_to_fail_with_tx_factory: Callable[..., None],
+ delete_df_collection: Callable[..., None],
):
"""
This test measures a complex Business situation, but then makes
@@ -985,11 +985,9 @@ class TestBusinessBalance:
delete_df_collection(coll=ledger_collection)
delete_df_collection(coll=task_adj_collection)
- from generalresearch.models.thl.user import User
-
- 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,
@@ -1002,22 +1000,17 @@ class TestBusinessBalance:
wall_req_cpi=Decimal("2.50"),
started=start + timedelta(days=2),
)
- payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
- bp_payout_factory(
+ brokerage_product_payout_event_factory(
product=u1.product,
amount=USDCent(250),
created=start + timedelta(days=3),
- skip_wallet_balance_check=True,
- 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),
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
)
session_with_tx_factory(
@@ -1042,73 +1035,80 @@ class TestBusinessBalance:
df = client_no_amm.compute(pop_ledger_merge.ddf(), sync=True)
assert df.shape == (20, 28)
- business.prebuild_balance(
- thl_pg_config=thl_web_rr,
- lm=lm,
+ gr_business.prebuild_balance(
+ product_manager=product_manager,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
)
- business.prebuild_balance(
- thl_pg_config=thl_web_rr,
- lm=lm,
+ gr_business.prebuild_balance(
+ product_manager=product_manager,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
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(
- thl_pg_config=thl_web_rr,
- lm=lm,
+ gr_business.prebuild_balance(
+ product_manager=product_manager,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
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(
- thl_pg_config=thl_web_rr,
- lm=lm,
+ gr_business.prebuild_balance(
+ product_manager=product_manager,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
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(
- thl_pg_config=thl_web_rr,
- lm=lm,
+ gr_business.prebuild_balance(
+ product_manager=product_manager,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
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(
- thl_pg_config=thl_web_rr,
- lm=lm,
+ gr_business.prebuild_balance(
+ product_manager=product_manager,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
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(
- thl_pg_config=thl_web_rr,
- lm=lm,
+ gr_business.prebuild_balance(
+ product_manager=product_manager,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
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)
+ assert isinstance(day3_bal, BusinessBalances)
+ assert isinstance(day4_bal, BusinessBalances)
+ assert isinstance(day5_bal, BusinessBalances)
+ assert isinstance(day6_bal, BusinessBalances)
assert day1_bal.payout == 238
assert day1_bal.retainer == 59
@@ -1136,9 +1136,8 @@ class TestBusinessBalance:
class TestBusinessMethods:
-
@pytest.fixture(scope="function")
- def start(self, utc_90days_ago) -> "datetime":
+ def start(self, utc_90days_ago: datetime) -> datetime:
s = utc_90days_ago.replace(microsecond=0)
return s
@@ -1149,72 +1148,74 @@ class TestBusinessMethods:
@pytest.fixture(scope="function")
def duration(
self,
- ) -> Optional["timedelta"]:
+ ) -> timedelta | None:
return None
- def test_cache_key(self, business, gr_redis):
- 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,
- business,
- gr_redis,
- gr_db,
- thl_web_rr,
- client_no_amm,
- mnt_filepath,
- lm,
- thl_lm,
+ gr_business: Business,
+ thl_web_rr: PostgresConfig,
+ client_no_amm: DaskClient,
+ mnt_filepath: GRLDatasets,
+ ledger_manager: LedgerManager,
+ thl_ledger_manager: ThlLedgerManager,
business_payout_event_manager,
- product_factory,
- membership_factory,
- team,
- session_with_tx_factory,
- user_factory,
+ gr_business_bank_account_manager: BusinessBankAccountManager,
+ product_manager: ProductManager,
+ product_factory: Callable[..., Product],
+ gr_team: Team,
+ session_with_tx_factory: Callable[..., Session],
+ user_factory: Callable[..., User],
ledger_collection,
- pop_ledger_merge,
- utc_60days_ago,
- delete_ledger_db,
- create_main_accounts,
- gr_redis_config,
- mnt_gr_api_dir,
+ pop_ledger_merge: PopLedgerMerge,
+ utc_60days_ago: datetime,
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ gr_redis_config: RedisConfig,
+ mnt_gr_api_dir: Path,
):
- assert gr_redis.get(name=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=business)
+ p1 = product_factory(team=gr_team, business=gr_business)
u1 = user_factory(product=p1)
# Business needs tx & incite to build balance
delete_ledger_db()
create_main_accounts()
- 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, started=utc_60days_ago)
ledger_collection.initial_load(client=None, sync=True)
pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection)
- business.set_cache(
- pg_config=gr_db,
+ gr_business.set_cache(
+ 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,
ds=mnt_filepath,
- lm=lm,
- thl_lm=thl_lm,
+ lm=ledger_manager,
+ thl_lm=thl_ledger_manager,
bpem=business_payout_event_manager,
pop_ledger=pop_ledger_merge,
mnt_gr_api=mnt_gr_api_dir,
)
- assert gr_redis.hgetall(name=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
# 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,
+ uuid=gr_business.uuid,
fields=[f"pop_financial:{year}"],
gr_redis_config=gr_redis_config,
)
@@ -1222,53 +1223,53 @@ class TestBusinessMethods:
def test_set_cache_business(
self,
- gr_user,
- business,
- gr_user_token,
- gr_redis,
- gr_db,
- thl_web_rr,
- product_factory,
- team,
- membership_factory,
- client_no_amm,
- mnt_filepath,
- lm,
- thl_lm,
+ gr_business: Business,
+ gr_db: PostgresConfig,
+ thl_web_rr: PostgresConfig,
+ product_factory: Callable[..., Product],
+ gr_team: Team,
+ client_no_amm: DaskClient,
+ mnt_filepath: GRLDatasets,
+ ledger_manager: LedgerManager,
+ thl_ledger_manager: ThlLedgerManager,
business_payout_event_manager,
- user_factory,
- delete_ledger_db,
- create_main_accounts,
- session_with_tx_factory,
+ product_manager: ProductManager,
+ gr_business_bank_account_manager: BusinessBankAccountManager,
+ user_factory: Callable[..., User],
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ session_with_tx_factory: Callable[..., Session],
ledger_collection,
- team_manager,
- pop_ledger_merge,
- gr_redis_config,
- utc_60days_ago,
- mnt_gr_api_dir,
+ gr_team_manager: TeamManager,
+ pop_ledger_merge: PopLedgerMerge,
+ gr_redis_config: RedisConfig,
+ utc_60days_ago: datetime,
+ mnt_gr_api_dir: Path,
):
from generalresearch.models.gr.business import Business
- p1 = product_factory(team=team, business=business)
+ p1 = product_factory(team=gr_team, business=gr_business)
u1 = user_factory(product=p1)
- team_manager.add_business(team=team, business=business)
+ gr_team_manager.add_business(team=gr_team, business=gr_business)
# Business needs tx & incite to build balance
delete_ledger_db()
create_main_accounts()
- 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, started=utc_60days_ago)
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(
+ 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,
client=client_no_amm,
ds=mnt_filepath,
- lm=lm,
- thl_lm=thl_lm,
+ lm=ledger_manager,
+ thl_lm=thl_ledger_manager,
bpem=business_payout_event_manager,
pop_ledger=pop_ledger_merge,
mnt_gr_api=mnt_gr_api_dir,
@@ -1276,7 +1277,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",
@@ -1295,11 +1296,16 @@ class TestBusinessMethods:
gr_redis_config=gr_redis_config,
)
- assert business.model_dump_json() == business2.model_dump_json()
+ assert isinstance(business2, Business)
+ 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)
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
assert business2.balance.balance == 48
assert business2.balance.net == 48
@@ -1312,39 +1318,39 @@ class TestBusinessMethods:
assert len(business2.bp_accounts) == 1
assert len(business2.bp_accounts) == len(business2.product_uuids)
+ assert isinstance(business2.pop_financial, list)
assert len(business2.pop_financial) == 1
assert business2.pop_financial[0].payout == business2.balance.payout
assert business2.pop_financial[0].net == business2.balance.net
def test_prebuild_enriched_session_parquet(
self,
- event_report_request,
enriched_session_merge,
- client_no_amm,
- wall_collection,
- session_collection,
- thl_web_rr,
- session_report_request,
- user_factory,
- start,
- session_factory,
- product_factory,
- delete_df_collection,
- business,
- mnt_filepath,
- mnt_gr_api_dir,
+ client_no_amm: DaskClient,
+ wall_collection: WallDFCollection,
+ product_manager: ProductManager,
+ 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],
+ gr_business: Business,
+ mnt_filepath: GRLDatasets,
+ mnt_gr_api_dir: Path,
):
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)
for i in range(50):
- s = session_factory(
+ session_factory(
user=u,
wall_count=1,
wall_req_cpi=Decimal("1.00"),
@@ -1360,8 +1366,8 @@ class TestBusinessMethods:
pg_config=thl_web_rr,
)
- business.prebuild_enriched_session_parquet(
- thl_pg_config=thl_web_rr,
+ gr_business.prebuild_enriched_session_parquet(
+ product_manager=product_manager,
ds=mnt_filepath,
client=client_no_amm,
mnt_gr_api=mnt_gr_api_dir,
@@ -1370,40 +1376,40 @@ 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)
def test_prebuild_enriched_wall_parquet(
self,
- event_report_request,
- enriched_session_merge,
enriched_wall_merge,
- client_no_amm,
- wall_collection,
- session_collection,
- thl_web_rr,
- session_report_request,
- user_factory,
- start,
- session_factory,
- product_factory,
- delete_df_collection,
- business,
- mnt_filepath,
- mnt_gr_api_dir,
+ client_no_amm: DaskClient,
+ wall_collection: WallDFCollection,
+ product_manager: ProductManager,
+ 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],
+ gr_business: Business,
+ mnt_filepath: GRLDatasets,
+ mnt_gr_api_dir: Path,
):
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)
for i in range(50):
- s = session_factory(
+ session_factory(
user=u,
wall_count=1,
wall_req_cpi=Decimal("1.00"),
@@ -1419,8 +1425,8 @@ class TestBusinessMethods:
pg_config=thl_web_rr,
)
- business.prebuild_enriched_wall_parquet(
- thl_pg_config=thl_web_rr,
+ gr_business.prebuild_enriched_wall_parquet(
+ product_manager=product_manager,
ds=mnt_filepath,
client=client_no_amm,
mnt_gr_api=mnt_gr_api_dir,
@@ -1429,6 +1435,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)
diff --git a/tests/models/gr/test_team.py b/tests/models/gr/test_team.py
index d728bbe..e853817 100644
--- a/tests/models/gr/test_team.py
+++ b/tests/models/gr/test_team.py
@@ -1,127 +1,180 @@
+from __future__ import annotations
+
import os
-from datetime import timedelta
+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
+from distributed.utils_test import (
+ client_no_amm,
+)
+
+from generalresearch.models.gr.business import Business
+from generalresearch.models.gr.team import Team
+from generalresearch.models.thl.product import Product
+
+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.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
+ from generalresearch.models.thl.user import User
+ from generalresearch.pg_helper import PostgresConfig
+ from generalresearch.redis_helper import RedisConfig
class TestTeam:
+ def test_init(self, gr_team: Team):
- def test_init(self, team):
- from generalresearch.models.gr.team import 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, gr_user_factory, gr_db):
- 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(gr_membership_manager=gr_membership_manager)
+ assert isinstance(gr_team.memberships, list)
+ assert len(gr_team.memberships) == 0
def test_memberships(
self,
- team,
- membership,
- gr_user,
- gr_user_factory,
- membership_factory,
- membership_manager,
- gr_db,
+ gr_team: Team,
+ gr_user: GRUser,
+ gr_membership,
+ gr_user_factory: Callable[..., GRUser],
+ 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(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
# 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(gr_membership_manager=gr_membership_manager)
+ assert len(gr_team.memberships) == 2
def test_gr_users(
- self, team, gr_user_factory, membership_manager, gr_db, gr_redis_config
+ self,
+ gr_team: Team,
+ gr_user_factory: Callable[..., GRUser],
+ gr_membership_manager: MembershipManager,
+ 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
+ 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=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
+ 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
- def test_businesses(self, team, business, team_manager, gr_db, gr_redis_config):
- from generalresearch.models.gr.business import Business
+ def test_businesses(
+ self,
+ gr_team: Team,
+ gr_business: Business,
+ team_manager: TeamManager,
+ 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(gr_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=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 == gr_business.uuid
- def test_products(self, team, product_factory, thl_web_rr):
- from generalresearch.models.thl.product import Product
+ def test_products(
+ self,
+ 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, gr_redis):
- 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,
- gr_redis,
- gr_db,
- thl_web_rr,
- gr_redis_config,
- client_no_amm,
- mnt_filepath,
- mnt_gr_api_dir,
- enriched_wall_merge,
- enriched_session_merge,
+ gr_team: Team,
+ gr_db: PostgresConfig,
+ thl_web_rr: PostgresConfig,
+ gr_redis_config: RedisConfig,
+ client_no_amm: DaskClient,
+ mnt_filepath: GRLDatasets,
+ 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,
):
- assert gr_redis.get(name=team.cache_key) is None
-
- team.set_cache(
- pg_config=gr_db,
- thl_web_rr=thl_web_rr,
+ client = gr_redis_config.create_redis_client()
+ assert client.get(name=gr_team.cache_key) is None
+
+ 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,
redis_config=gr_redis_config,
client=client_no_amm,
ds=mnt_filepath,
@@ -130,33 +183,36 @@ class TestTeamMethods:
enriched_session=enriched_session_merge,
)
- assert gr_redis.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,
- gr_user,
- gr_user_token,
- gr_redis,
- gr_db,
- thl_web_rr,
- product_factory,
- team,
- membership_factory,
- gr_redis_config,
- client_no_amm,
- mnt_filepath,
- mnt_gr_api_dir,
- enriched_wall_merge,
- enriched_session_merge,
+ gr_user: GRUser,
+ gr_db: PostgresConfig,
+ thl_web_rr: PostgresConfig,
+ product_factory: Callable[..., Product],
+ gr_team: Team,
+ gr_membership_factory: Callable[..., Membership],
+ gr_redis_config: RedisConfig,
+ mnt_filepath: GRLDatasets,
+ 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
- p1 = product_factory(team=team)
- membership_factory(team=team, gr_user=gr_user)
+ p1 = product_factory(team=gr_team)
+ gr_membership_factory(gr_team=gr_team, gr_user=gr_user)
- team.set_cache(
- pg_config=gr_db,
- thl_web_rr=thl_web_rr,
+ 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,
redis_config=gr_redis_config,
client=client_no_amm,
ds=mnt_filepath,
@@ -166,46 +222,47 @@ 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,
)
- assert team.model_dump_json() == team2.model_dump_json()
+ assert isinstance(team2, Team)
+ assert isinstance(team2.products, list)
+ assert isinstance(team2.gr_users, list)
+ 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]
def test_prebuild_enriched_session_parquet(
self,
- event_report_request,
- enriched_session_merge,
- client_no_amm,
- wall_collection,
- session_collection,
- thl_web_rr,
- session_report_request,
- user_factory,
- start,
- session_factory,
- product_factory,
- delete_df_collection,
- business,
- mnt_filepath,
- mnt_gr_api_dir,
- team,
+ enriched_session_merge: EnrichedSessionMerge,
+ client_no_amm: DaskClient,
+ 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],
+ mnt_filepath: GRLDatasets,
+ mnt_gr_api_dir: Path,
+ gr_team: Team,
+ product_manager: ProductManager,
):
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)
for i in range(50):
- s = session_factory(
+ session_factory(
user=u,
wall_count=1,
wall_req_cpi=Decimal("1.00"),
@@ -221,8 +278,8 @@ class TestTeamMethods:
pg_config=thl_web_rr,
)
- team.prebuild_enriched_session_parquet(
- thl_pg_config=thl_web_rr,
+ gr_team.prebuild_enriched_session_parquet(
+ product_manager=product_manager,
ds=mnt_filepath,
client=client_no_amm,
mnt_gr_api=mnt_gr_api_dir,
@@ -231,41 +288,38 @@ 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)
def test_prebuild_enriched_wall_parquet(
self,
- event_report_request,
- enriched_session_merge,
- enriched_wall_merge,
- client_no_amm,
- wall_collection,
- session_collection,
- thl_web_rr,
- session_report_request,
- user_factory,
- start,
- session_factory,
- product_factory,
- delete_df_collection,
- business,
- mnt_filepath,
- mnt_gr_api_dir,
- team,
+ enriched_wall_merge: EnrichedWallMerge,
+ client_no_amm: DaskClient,
+ wall_collection: WallDFCollection,
+ session_collection: EnrichedSessionMerge,
+ thl_web_rr: PostgresConfig,
+ user_factory: Callable[..., User],
+ start: datetime,
+ session_factory: Callable[..., Session],
+ product_factory: Callable[..., Product],
+ delete_df_collection: Callable[..., None],
+ mnt_filepath: GRLDatasets,
+ mnt_gr_api_dir: Path,
+ gr_team: Team,
+ product_manager: ProductManager,
):
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)
for i in range(50):
- s = session_factory(
+ session_factory(
user=u,
wall_count=1,
wall_req_cpi=Decimal("1.00"),
@@ -281,8 +335,8 @@ class TestTeamMethods:
pg_config=thl_web_rr,
)
- team.prebuild_enriched_wall_parquet(
- thl_pg_config=thl_web_rr,
+ gr_team.prebuild_enriched_wall_parquet(
+ product_manager=product_manager,
ds=mnt_filepath,
client=client_no_amm,
mnt_gr_api=mnt_gr_api_dir,
@@ -291,6 +345,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)
diff --git a/tests/models/innovate/test_question.py b/tests/models/innovate/test_question.py
index 330f919..ea2fc8c 100644
--- a/tests/models/innovate/test_question.py
+++ b/tests/models/innovate/test_question.py
@@ -1,15 +1,17 @@
-from generalresearch.models import Source
+from __future__ import annotations
+
+from generalresearch.models.definitions 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_offerwall_parse_response.py b/tests/models/legacy/test_offerwall_parse_response.py
index b1c96ad..93f5c26 100644
--- a/tests/models/legacy/test_offerwall_parse_response.py
+++ b/tests/models/legacy/test_offerwall_parse_response.py
@@ -1,6 +1,8 @@
+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_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 224334a..f14c1a7 100644
--- a/tests/models/legacy/test_user_question_answer_in.py
+++ b/tests/models/legacy/test_user_question_answer_in.py
@@ -1,9 +1,25 @@
+from __future__ import annotations
+
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.models.definitions import Source
+from generalresearch.models.legacy.questions import (
+ UserQuestionAnswers,
+)
+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
@@ -15,21 +31,11 @@ class TestUserQuestionAnswers:
def test_json_init(
self,
- product_manager,
- user_manager,
- session_manager,
- wall_manager,
- user_factory,
- product,
- session_factory,
- utc_hour_ago,
+ user_factory: Callable[..., User],
+ product: Product,
+ 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)
@@ -60,11 +66,8 @@ class TestUserQuestionAnswers:
assert isinstance(instance, UserQuestionAnswers)
def test_simple_validation_errors(
- self, product_manager, user_manager, session_manager, wall_manager
+ self,
):
- from generalresearch.models.legacy.questions import (
- UserQuestionAnswers,
- )
with pytest.raises(ValueError):
UserQuestionAnswers.model_validate(
@@ -114,7 +117,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(
{
@@ -139,9 +142,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:
@@ -161,11 +161,11 @@ class TestUserQuestionAnswers:
def test_allow_answer_failures_silent(
self,
- user_manager,
- product,
- user_factory,
- utc_hour_ago,
- session_factory,
+ user_manager: UserManager,
+ product: Product,
+ user_factory: Callable[..., User],
+ utc_hour_ago: datetime,
+ session_factory: Callable[..., Session],
):
"""
There are many instances where suppliers may be submitting answers
@@ -173,11 +173,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)
@@ -263,12 +258,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(
@@ -282,7 +277,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}
@@ -294,8 +289,8 @@ class TestUserQuestionAnswerIn:
UserQuestionAnswerIn,
)
- answer = ["aaa" for i in range(5)]
- with pytest.raises(ValueError) as cm:
+ 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 bedf9c2..c1141fb 100644
--- a/tests/models/morning/test.py
+++ b/tests/models/morning/test.py
@@ -1,4 +1,6 @@
-from datetime import datetime, timezone
+from __future__ import annotations
+
+from datetime import UTC, datetime
from generalresearch.models.morning.question import MorningQuestion
@@ -163,8 +165,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/network/__init__.py b/tests/models/network/__init__.py
deleted file mode 100644
index e69de29..0000000
--- a/tests/models/network/__init__.py
+++ /dev/null
diff --git a/tests/models/network/test_mtr.py b/tests/models/network/test_mtr.py
deleted file mode 100644
index 2965300..0000000
--- a/tests/models/network/test_mtr.py
+++ /dev/null
@@ -1,26 +0,0 @@
-from generalresearch.models.network.mtr.execute import execute_mtr
-import faker
-
-from generalresearch.models.network.tool_run import ToolName, ToolClass
-
-fake = faker.Faker()
-
-
-def test_execute_mtr(toolrun_manager):
- 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 a135a13..0000000
--- a/tests/models/network/test_nmap.py
+++ /dev/null
@@ -1,30 +0,0 @@
-import subprocess
-
-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
-
-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 = 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
- 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 abc83c9..0000000
--- a/tests/models/network/test_nmap_parser.py
+++ /dev/null
@@ -1,22 +0,0 @@
-import os
-
-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")
- with open(fp) as f:
- data = f.read()
- return data
-
-
-def test_nmap_xml_parser(nmap_raw_output, nmap_raw_output_2):
- n = parse_nmap_xml(nmap_raw_output)
- assert n.tcp_open_ports == [61232]
- 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 5c3b024..0000000
--- a/tests/models/network/test_rdns.py
+++ /dev/null
@@ -1,34 +0,0 @@
-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
-
-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)
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 68d7838..10ce884 100644
--- a/tests/models/prodege/test_survey_participation.py
+++ b/tests/models/prodege/test_survey_participation.py
@@ -1,16 +1,19 @@
-from datetime import datetime, timedelta, timezone
+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=timezone.utc)
+ now = datetime.now(tz=UTC)
pp = ProdegePastParticipation.from_api(
{
"participation_project_ids": [152677146, 152803285],
@@ -84,12 +87,8 @@ class TestProdegeParticipation:
assert not pp.is_eligible(upps)
def test_include(self):
- from generalresearch.models.prodege.survey import (
- ProdegePastParticipation,
- 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..d469530 100644
--- a/tests/models/spectrum/test_question.py
+++ b/tests/models/spectrum/test_question.py
@@ -1,17 +1,19 @@
-from datetime import datetime, timezone
+from __future__ import annotations
-from generalresearch.models import Source
+from datetime import UTC, datetime
+
+from generalresearch.models.definitions 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,
)
@@ -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",
@@ -43,7 +46,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,
@@ -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",
@@ -85,7 +90,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 +165,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..02c5d3f 100644
--- a/tests/models/spectrum/test_survey.py
+++ b/tests/models/spectrum/test_survey.py
@@ -1,15 +1,25 @@
-from datetime import datetime, timezone
+from __future__ import annotations
+
+from datetime import UTC, datetime
from decimal import Decimal
+from generalresearch.models.definitions 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,29 +122,17 @@ 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 = {
"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,
@@ -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"],
@@ -212,7 +202,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 +230,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 +245,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,
@@ -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
@@ -318,11 +310,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,
@@ -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"]),
@@ -411,3 +405,15 @@ class TestSpectrumSurvey:
assert (None, {"c", "d"}) == s.determine_eligibility_soft(
{"a": True, "b": True, "c": None, "d": None}
)
+
+
+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
+ assert c3.criterion_hash in survey.qualifications
diff --git a/tests/models/spectrum/test_survey_manager.py b/tests/models/spectrum/test_survey_manager.py
index 582093c..0300956 100644
--- a/tests/models/spectrum/test_survey_manager.py
+++ b/tests/models/spectrum/test_survey_manager.py
@@ -1,72 +1,36 @@
-import copy
+from __future__ import annotations
+
import logging
-from datetime import timezone, datetime
+from datetime import UTC, datetime
from decimal import Decimal
+from typing import TYPE_CHECKING, Any
from pymysql import IntegrityError
+from generalresearch.config import is_debug
-logger = logging.getLogger()
+if TYPE_CHECKING:
+ 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=timezone.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),
- "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=timezone.utc)
+ now = datetime.now(tz=UTC)
spectrum_rw.execute_sql_query(
query=f"""
DELETE FROM `{spectrum_rw.db}`.spectrum_survey
@@ -74,26 +38,31 @@ 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=timezone.utc)
+ now = datetime.now(tz=UTC)
spectrum_rw.execute_sql_query(
query=f"""
DELETE FROM `{spectrum_rw.db}`.spectrum_survey
@@ -101,14 +70,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
@@ -123,8 +91,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 40cff88..e946126 100644
--- a/tests/models/test_currency.py
+++ b/tests/models/test_currency.py
@@ -3,27 +3,29 @@ 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
import pytest
+from generalresearch.currency import USDCent, USDMill, format_usd_cent
+
class TestUSDCentModel:
def test_construct_int(self):
- from generalresearch.currency import USDCent
- for i in range(100):
+ for _ in range(100):
int_val = randint(0, 999_999)
instance = USDCent(int_val)
assert int_val == instance
def test_construct_float(self):
- from generalresearch.currency import USDCent
+ float_val: float = 10.6789
with pytest.warns(expected_warning=Warning) as record:
- float_val: float = 10.6789
instance = USDCent(float_val)
assert len(record) == 1
@@ -34,10 +36,9 @@ class TestUSDCentModel:
assert instance == 10
def test_construct_decimal(self):
- from generalresearch.currency import USDCent
+ decimal_val: Decimal = Decimal("10.0")
with pytest.warns(expected_warning=Warning) as record:
- decimal_val: Decimal = Decimal("10.0")
instance = USDCent(decimal_val)
assert len(record) == 1
@@ -50,8 +51,8 @@ class TestUSDCentModel:
assert instance == 10
# Now with rounding
+ decimal_val: Decimal = Decimal("10.6789")
with pytest.warns(Warning) as record:
- decimal_val: Decimal = Decimal("10.6789")
instance = USDCent(decimal_val)
assert len(record) == 1
@@ -64,16 +65,12 @@ class TestUSDCentModel:
assert instance == 10
def test_construct_negative(self):
- from generalresearch.currency import USDCent
-
with pytest.raises(expected_exception=ValueError) as cm:
USDCent(-1)
assert "USDCent not be less than zero" in str(cm.value)
def test_operation_add(self):
- from generalresearch.currency import USDCent
-
- for i in range(100):
+ for _ in range(100):
int_val1 = randint(0, 999_999)
int_val2 = randint(0, 999_999)
@@ -83,9 +80,7 @@ class TestUSDCentModel:
assert int_val1 + int_val2 == instance1 + instance2
def test_operation_subtract(self):
- from generalresearch.currency import USDCent
-
- for i in range(100):
+ for _ in range(100):
int_val1 = randint(500_000, 999_999)
int_val2 = randint(0, 499_999)
@@ -95,21 +90,17 @@ class TestUSDCentModel:
assert int_val1 - int_val2 == instance1 - instance2
def test_operation_subtract_to_neg(self):
- from generalresearch.currency import USDCent
-
- for i in range(100):
+ for _ in range(100):
int_val = randint(0, 999_999)
instance = USDCent(int_val)
with pytest.raises(expected_exception=ValueError) as cm:
- instance - USDCent(1_000_000)
+ _ = instance - USDCent(1_000_000)
assert "USDCent not be less than zero" in str(cm.value)
def test_operation_multiply(self):
- from generalresearch.currency import USDCent
-
- for i in range(100):
+ for _ in range(100):
int_val1 = randint(0, 999_999)
int_val2 = randint(0, 999_999)
@@ -119,15 +110,11 @@ class TestUSDCentModel:
assert int_val1 * int_val2 == instance1 * instance2
def test_operation_div(self):
- from generalresearch.currency import USDCent
-
with pytest.raises(ValueError) as cm:
- USDCent(10) / 2
+ _ = USDCent(10) / 2
assert "Division not allowed for USDCent" in str(cm.value)
def test_operation_result_type(self):
- from generalresearch.currency import USDCent
-
int_val = randint(1, 999_999)
instance = USDCent(int_val)
@@ -141,36 +128,30 @@ class TestUSDCentModel:
assert isinstance(res_multipy, USDCent)
def test_operation_partner_add(self):
- from generalresearch.currency import USDCent
-
int_val = randint(1, 999_999)
instance = USDCent(int_val)
with pytest.raises(expected_exception=AssertionError):
- instance + 0.10
+ _ = instance + 0.10
with pytest.raises(expected_exception=AssertionError):
- instance + Decimal(".10")
+ _ = instance + Decimal(".10")
with pytest.raises(expected_exception=AssertionError):
- instance + "9.9"
+ _ = instance + "9.9"
with pytest.raises(expected_exception=AssertionError):
- instance + True
+ _ = instance + True
def test_abs(self):
- from generalresearch.currency import USDCent
-
- for i in range(100):
+ for _ in range(100):
int_val = abs(randint(0, 999_999))
instance = abs(USDCent(int_val))
assert int_val == instance
def test_str(self):
- from generalresearch.currency import USDCent
-
- for i in range(100):
+ for _ in range(100):
int_val = randint(0, 999_999)
instance = USDCent(int_val)
@@ -180,8 +161,6 @@ class TestUSDCentModel:
"""There is no correct answer here, but we at least want to make sure
that a USDCent is returned
"""
- from generalresearch.currency import USDCent
-
res = USDCent(10) // 1.2
assert not isinstance(res, USDCent)
assert isinstance(res, float)
@@ -206,18 +185,14 @@ class TestUSDCentModel:
class TestUSDMillModel:
def test_construct_int(self):
- from generalresearch.currency import USDMill
-
- for i in range(100):
+ for _ in range(100):
int_val = randint(0, 999_999)
instance = USDMill(int_val)
assert int_val == instance
def test_construct_float(self):
- from generalresearch.currency import USDMill
-
+ float_val: float = 10.6789
with pytest.warns(expected_warning=Warning) as record:
- float_val: float = 10.6789
instance = USDMill(float_val)
assert len(record) == 1
@@ -228,10 +203,8 @@ class TestUSDMillModel:
assert instance == 10
def test_construct_decimal(self):
- from generalresearch.currency import USDMill
-
+ decimal_val: Decimal = Decimal("10.0")
with pytest.warns(expected_warning=Warning) as record:
- decimal_val: Decimal = Decimal("10.0")
instance = USDMill(decimal_val)
assert len(record) == 1
@@ -244,10 +217,11 @@ class TestUSDMillModel:
assert instance == 10
# Now with rounding
+ decimal_val: Decimal = Decimal("10.6789")
with pytest.warns(expected_warning=Warning) as record:
- decimal_val: Decimal = Decimal("10.6789")
instance = USDMill(decimal_val)
+ assert isinstance(instance, USDMill)
assert len(record) == 1
assert (
"USDMill init with a Decimal. Rounding behavior may be unexpected"
@@ -258,16 +232,12 @@ class TestUSDMillModel:
assert instance == 10
def test_construct_negative(self):
- from generalresearch.currency import USDMill
-
with pytest.raises(expected_exception=ValueError) as cm:
USDMill(-1)
assert "USDMill not be less than zero" in str(cm.value)
def test_operation_add(self):
- from generalresearch.currency import USDMill
-
- for i in range(100):
+ for _ in range(100):
int_val1 = randint(0, 999_999)
int_val2 = randint(0, 999_999)
@@ -277,9 +247,7 @@ class TestUSDMillModel:
assert int_val1 + int_val2 == instance1 + instance2
def test_operation_subtract(self):
- from generalresearch.currency import USDMill
-
- for i in range(100):
+ for _ in range(100):
int_val1 = randint(500_000, 999_999)
int_val2 = randint(0, 499_999)
@@ -289,21 +257,17 @@ class TestUSDMillModel:
assert int_val1 - int_val2 == instance1 - instance2
def test_operation_subtract_to_neg(self):
- from generalresearch.currency import USDMill
-
- for i in range(100):
+ for _ in range(100):
int_val = randint(0, 999_999)
instance = USDMill(int_val)
with pytest.raises(expected_exception=ValueError) as cm:
- instance - USDMill(1_000_000)
+ _ = instance - USDMill(1_000_000)
assert "USDMill not be less than zero" in str(cm.value)
def test_operation_multiply(self):
- from generalresearch.currency import USDMill
-
- for i in range(100):
+ for _ in range(100):
int_val1 = randint(0, 999_999)
int_val2 = randint(0, 999_999)
@@ -313,15 +277,11 @@ class TestUSDMillModel:
assert int_val1 * int_val2 == instance1 * instance2
def test_operation_div(self):
- from generalresearch.currency import USDMill
-
with pytest.raises(ValueError) as cm:
- USDMill(10) / 2
+ _ = USDMill(10) / 2
assert "Division not allowed for USDMill" in str(cm.value)
def test_operation_result_type(self):
- from generalresearch.currency import USDMill
-
int_val = randint(1, 999_999)
instance = USDMill(int_val)
@@ -335,36 +295,30 @@ class TestUSDMillModel:
assert isinstance(res_multipy, USDMill)
def test_operation_partner_add(self):
- from generalresearch.currency import USDMill
-
int_val = randint(1, 999_999)
instance = USDMill(int_val)
with pytest.raises(expected_exception=AssertionError):
- instance + 0.10
+ _ = instance + 0.10
with pytest.raises(expected_exception=AssertionError):
- instance + Decimal(".10")
+ _ = instance + Decimal(".10")
with pytest.raises(expected_exception=AssertionError):
- instance + "9.9"
+ _ = instance + "9.9"
with pytest.raises(expected_exception=AssertionError):
- instance + True
+ _ = instance + True
def test_abs(self):
- from generalresearch.currency import USDMill
-
- for i in range(100):
+ for _ in range(100):
int_val = abs(randint(0, 999_999))
instance = abs(USDMill(int_val))
assert int_val == instance
def test_str(self):
- from generalresearch.currency import USDMill
-
- for i in range(100):
+ for _ in range(100):
int_val = randint(0, 999_999)
instance = USDMill(int_val)
@@ -374,8 +328,6 @@ class TestUSDMillModel:
"""There is no correct answer here, but we at least want to make sure
that a USDMill is returned
"""
- from generalresearch.currency import USDCent, USDMill
-
res = USDMill(10) // 1.2
assert not isinstance(res, USDMill)
assert isinstance(res, float)
@@ -400,11 +352,7 @@ class TestUSDMillModel:
class TestNegativeFormatting:
def test_pos(self):
- from generalresearch.currency import format_usd_cent
-
assert "-$987.65" == format_usd_cent(-98765)
def test_neg(self):
- from generalresearch.currency import format_usd_cent
-
assert "-$123.45" == format_usd_cent(-12345)
diff --git a/tests/models/test_device.py b/tests/models/test_device.py
index bf72c81..fdbd906 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.definitions 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 bd548b3..a1da961 100644
--- a/tests/models/test_finance.py
+++ b/tests/models/test_finance.py
@@ -1,44 +1,43 @@
-from datetime import datetime, timedelta, timezone
+from __future__ import annotations
+
+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 Callable
+from typing import TYPE_CHECKING
from uuid import uuid4
+import dask.dataframe as dd
import pandas as pd
import pytest
from dask.distributed import Client as DaskClient
# noinspection PyUnresolvedReferences
-from distributed.utils_test import (
- client_no_amm,
-)
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.product import Product
-from generalresearch.models.thl.session import Session
-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,
-)
+
+if TYPE_CHECKING:
+ from generalresearch.incite.collections.thl_web import LedgerDFCollection
+ from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge
+ from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager
+ from generalresearch.managers.thl.product import ProductManager
+ from generalresearch.models.thl.ledger import LedgerAccount
+ from generalresearch.models.thl.product import Product
+ from generalresearch.models.thl.session import Session
+ from generalresearch.models.thl.user import User
fake = Faker()
class TestProductBalanceInitialize:
-
def test_unknown_fields(self):
with pytest.raises(expected_exception=ValueError):
ProductBalances.model_validate(
@@ -210,6 +209,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
@@ -244,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}
@@ -653,37 +653,37 @@ class TestBusinessBalanceInitialize:
@pytest.mark.parametrize(
- argnames="offset, duration",
+ argnames="duration",
argvalues=list(
iter_product(
- ["12h", "2D"],
[timedelta(days=2), timedelta(days=5)],
)
),
)
class TestProductFinanceData:
-
def test_base(
self,
+ ledger_collection: LedgerDFCollection,
+ pop_ledger_merge,
+ client_no_amm,
+ duration: timedelta,
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, 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,
end_date=item.finish,
- tzinfo=timezone.utc,
+ tzinfo=UTC,
)
session_with_tx_factory(started=rand_item_time, user=u)
@@ -697,10 +697,9 @@ class TestProductFinanceData:
item_finishes = [i.finish for i in ledger_collection.items]
item_finishes.sort(reverse=True)
- last_item_finish = item_finishes[0]
# --
- 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,
@@ -732,17 +731,7 @@ class TestProductFinanceData:
assert len(res) == len({i.time for i in res})
-@pytest.mark.parametrize(
- argnames="offset, duration",
- argvalues=list(
- iter_product(
- ["12h", "2D"],
- [timedelta(days=2), timedelta(days=5)],
- )
- ),
-)
class TestPOPFinancialData:
-
def test_base(
self,
client_no_amm: DaskClient,
@@ -751,19 +740,16 @@ class TestPOPFinancialData:
user_factory: Callable[..., User],
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,
+ thl_ledger_manager: ThlLedgerManager,
+ delete_df_collection: Callable[..., None],
+ delete_ledger_db: Callable[..., None],
):
# -- Build & Setup
delete_ledger_db()
create_main_accounts()
delete_df_collection(coll=ledger_collection)
- # assert ledger_collection.start is None
- # assert ledger_collection.offset is None
users = []
for _ in range(5):
@@ -773,7 +759,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)
@@ -792,8 +778,10 @@ class TestPOPFinancialData:
last_item_finish = item_finishes[0]
accounts = []
- for user 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]
@@ -809,6 +797,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()
@@ -821,25 +810,13 @@ 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
- assert isinstance(duration, timedelta)
-
# -- Teardown
delete_df_collection(ledger_collection)
-@pytest.mark.parametrize(
- argnames="offset, duration",
- argvalues=list(
- iter_product(
- ["12h", "1D"],
- [timedelta(days=2), timedelta(days=3)],
- )
- ),
-)
class TestBusinessBalanceData:
def test_from_pandas(
self,
@@ -848,15 +825,14 @@ class TestBusinessBalanceData:
pop_ledger_merge: PopLedgerMerge,
user_factory: Callable[..., User],
product: Product,
- create_main_accounts,
- thl_lm: ThlLedgerManager,
- thl_web_rr,
- delete_df_collection,
- delete_ledger_db,
+ create_main_accounts: Callable[..., None],
+ thl_ledger_manager: ThlLedgerManager,
+ product_manager: ProductManager,
+ 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()
@@ -870,7 +846,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)
@@ -880,7 +856,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,
@@ -888,15 +866,18 @@ class TestBusinessBalanceData:
columns=numerical_col_names + ["account_id"],
filters=[("account_id", "in", [account.uuid])],
)
+ assert isinstance(ddf, dd.DataFrame)
ddf = ddf.groupby("account_id").sum()
df: pd.DataFrame = client_no_amm.compute(collections=ddf, sync=True)
assert isinstance(df, pd.DataFrame)
instance = BusinessBalances.from_pandas(
- input_data=df, accounts=[account], thl_pg_config=thl_web_rr
+ product_manager=product_manager,
+ input_data=df,
+ accounts=[account],
)
- balance: int = thl_lm.get_account_balance(account=account)
+ balance: int = thl_ledger_manager.get_account_balance(account=account)
assert instance.balance == balance
assert instance.net == balance
diff --git a/tests/models/thl/question/test_question_info.py b/tests/models/thl/question/test_question_info.py
index 945ee7a..af8d2b9 100644
--- a/tests/models/thl/question/test_question_info.py
+++ b/tests/models/thl/question/test_question_info.py
@@ -1,145 +1,16 @@
+from __future__ import annotations
+
from generalresearch.models.thl.profiling.upk_property import (
- UpkProperty,
ProfilingInfo,
+ UpkProperty,
)
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 27091bb..5d8605a 100644
--- a/tests/models/thl/test_adjustments.py
+++ b/tests/models/thl/test_adjustments.py
@@ -1,36 +1,42 @@
-from datetime import datetime, timedelta, timezone
+from __future__ import annotations
+
+from collections.abc import Callable
+from datetime import UTC, datetime, timedelta
from decimal import Decimal
-from typing import Callable
+from typing import TYPE_CHECKING
import pytest
-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,
+ Session,
+ Wall,
)
-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)
+if TYPE_CHECKING:
+ from generalresearch.managers.thl.session import SessionManager
+ from generalresearch.managers.thl.wall import WallManager
+ from generalresearch.models.thl.user import User
+
+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:
-
@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 +45,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))
@@ -48,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(
@@ -60,7 +65,7 @@ class TestSessionAdjustments:
)
# Confirm only the last Wall Event is a complete
- assert not s1.wall_events[0].status == Status.COMPLETE
+ assert s1.wall_events[0].status != Status.COMPLETE
assert s1.wall_events[1].status == Status.COMPLETE
# Confirm the Session is marked as finished and the simple brokerage
@@ -71,9 +76,11 @@ 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,14 +432,14 @@ 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(
user=user,
wall_count=1,
- wall_req_cpi=Decimal("1"),
+ wall_req_cpi=Decimal(1),
final_status=Status.COMPLETE,
started=utc_hour_ago,
)
@@ -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(
@@ -525,22 +539,20 @@ 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,
)
w1 = s1.wall_events[0]
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()
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,
@@ -562,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
@@ -590,15 +603,14 @@ 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,
- "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.
@@ -628,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(
@@ -644,15 +659,14 @@ 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,
- "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.
@@ -664,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
@@ -708,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..ef97166 100644
--- a/tests/models/thl/test_buyer.py
+++ b/tests/models/thl/test_buyer.py
@@ -1,4 +1,6 @@
-from generalresearch.models import Source
+from __future__ import annotations
+
+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 0fbd4cc..ed8477b 100644
--- a/tests/models/thl/test_contest/test_contest.py
+++ b/tests/models/thl/test_contest/test_contest.py
@@ -1,9 +1,13 @@
-from typing import Callable
+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 8b714ee..a639261 100644
--- a/tests/models/thl/test_contest/test_leaderboard_contest.py
+++ b/tests/models/thl/test_contest/test_leaderboard_contest.py
@@ -1,7 +1,11 @@
-from datetime import timezone
+from __future__ import annotations
+
+from datetime import UTC
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
+from redis import Redis
from generalresearch.currency import USDCent
from generalresearch.managers.leaderboard.manager import LeaderboardManager
@@ -17,16 +21,20 @@ 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):
@pytest.fixture
def leaderboard_contest(
- self, product: Product, thl_redis, user_manager
- ) -> "LeaderboardContest":
+ self, product: Product, thl_redis_client: Redis, user_manager: UserManager
+ ) -> LeaderboardContest:
board_key = f"leaderboard:{product.uuid}:us:weekly:2025-05-26:complete_count"
c = LeaderboardContest(
@@ -59,16 +67,22 @@ 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, thl_redis, user_1, user_2):
+ def test_init(
+ self,
+ leaderboard_contest: LeaderboardContest,
+ thl_redis_client: Redis,
+ user_1: User,
+ user_2: User,
+ ):
model = leaderboard_contest.leaderboard_model
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,
@@ -83,15 +97,22 @@ 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_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,
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)
@@ -102,10 +123,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 d7920f0..e71851e 100644
--- a/tests/models/thl/test_contest/test_raffle_contest.py
+++ b/tests/models/thl/test_contest/test_raffle_contest.py
@@ -1,4 +1,8 @@
+from __future__ import annotations
+
from collections import Counter
+from datetime import datetime
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
@@ -18,9 +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 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):
@@ -42,7 +49,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 +64,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 +87,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 +133,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 +171,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 +210,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 +237,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 +264,12 @@ 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: 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 257de3c..7c48dbd 100644
--- a/tests/models/thl/test_ledger.py
+++ b/tests/models/thl/test_ledger.py
@@ -1,4 +1,6 @@
-from datetime import datetime, timezone
+from __future__ import annotations
+
+from datetime import UTC, datetime
from uuid import uuid4
import pytest
@@ -21,7 +23,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_marketplace_condition.py b/tests/models/thl/test_marketplace_condition.py
index 8a4b25c..6936a7c 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.definitions 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(
@@ -137,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,
@@ -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(
@@ -265,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,
@@ -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 3a51328..cc00f33 100644
--- a/tests/models/thl/test_payout.py
+++ b/tests/models/thl/test_payout.py
@@ -1,10 +1,122 @@
+from __future__ import annotations
+
+from uuid import uuid4
+
+import pytest
+from pydantic import ValidationError
+
+from generalresearch.currency import USDCent
+from generalresearch.models.gr import Team
+from generalresearch.models.gr.business import (
+ Business,
+ BusinessAddress,
+)
+from generalresearch.models.gr.definitions import BusinessType
+from generalresearch.models.thl.payout import (
+ BrokerageProductPayoutEvent,
+ BusinessPayoutEvent,
+)
+from generalresearch.models.thl.wallet.definitions import PayoutType
+
+
class TestBusinessPayoutEvent:
def test_validate(self):
- from generalresearch.models.gr.business import Business
- instance = Business.model_validate_json(
- json_data='{"id":123,"uuid":"947f6ba5250d442b9a66cde9ee33605a","name":"Example » Demo","kind":"c","tax_number":null,"contact":null,"addresses":[],"teams":[{"id":53,"uuid":"8e4197dcaefe4f1f831a02b212e6b44a","name":"Example » Demo","memberships":null,"gr_users":null,"businesses":null,"products":null}],"products":[{"id":"fc23e741b5004581b30e6478363525df","id_int":1234,"name":"Example","enabled":true,"payments_enabled":true,"created":"2025-04-14T13:25:37.279403Z","team_id":"9e4197dcaefe4f1f831a02b212e6b44a","business_id":"857f6ba6160d442b9a66cde9ee33605a","tags":[],"commission_pct":"0.050000","redirect_url":"https://pam-api-us.reppublika.com/v2/public/4970ef00-0ef7-11f0-9962-05cb6323c84c/grl/status","harmonizer_domain":"https://talk.generalresearch.com/","sources_config":{"user_defined":[{"name":"w","active":false,"banned_countries":[],"allow_mobile_ip":true,"supplier_id":null,"allow_pii_only_buyers":false,"allow_unhashed_buyers":false,"withhold_profiling":false,"pass_unconditional_eligible_unknowns":true,"address":null,"allow_vpn":null,"distribute_harmonizer_active":null}]},"session_config":{"max_session_len":600,"max_session_hard_retry":5,"min_payout":"0.14"},"payout_config":{"payout_format":null,"payout_transformation":null},"user_wallet_config":{"enabled":false,"amt":false,"supported_payout_types":["CASH_IN_MAIL","PAYPAL","TANGO"],"min_cashout":null},"user_create_config":{"min_hourly_create_limit":0,"max_hourly_create_limit":null},"offerwall_config":{},"profiling_config":{"enabled":true,"grs_enabled":true,"n_questions":null,"max_questions":10,"avg_question_count":5.0,"task_injection_freq_mult":1.0,"non_us_mult":2.0,"hidden_questions_expiration_hours":168},"user_health_config":{"banned_countries":[],"allow_ban_iphist":true},"yield_man_config":{},"balance":null,"payouts_total_str":null,"payouts_total":null,"payouts":null,"user_wallet":{"enabled":false,"amt":false,"supported_payout_types":["CASH_IN_MAIL","PAYPAL","TANGO"],"min_cashout":null}}],"bank_accounts":[],"balance":{"product_balances":[{"product_id":"fc14e741b5004581b30e6478363414df","last_event":null,"bp_payment_credit":780251,"adjustment_credit":4678,"adjustment_debit":26446,"supplier_credit":0,"supplier_debit":451513,"user_bonus_credit":0,"user_bonus_debit":0,"issued_payment":0,"payout":780251,"payout_usd_str":"$7,802.51","adjustment":-21768,"expense":0,"net":758483,"payment":451513,"payment_usd_str":"$4,515.13","balance":306970,"retainer":76742,"retainer_usd_str":"$767.42","available_balance":230228,"available_balance_usd_str":"$2,302.28","recoup":0,"recoup_usd_str":"$0.00","adjustment_percent":0.027898714644390074}],"payout":780251,"payout_usd_str":"$7,802.51","adjustment":-21768,"expense":0,"net":758483,"net_usd_str":"$7,584.83","payment":451513,"payment_usd_str":"$4,515.13","balance":306970,"balance_usd_str":"$3,069.70","retainer":76742,"retainer_usd_str":"$767.42","available_balance":230228,"available_balance_usd_str":"$2,302.28","adjustment_percent":0.027898714644390074,"recoup":0,"recoup_usd_str":"$0.00"},"payouts_total_str":"$4,515.13","payouts_total":451513,"payouts":[{"bp_payouts":[{"uuid":"40cf2c3c341e4f9d985be4bca43e6116","debit_account_uuid":"3a058056da85493f9b7cdfe375aad0e0","cashout_method_uuid":"602113e330cf43ae85c07d94b5100291","created":"2025-08-02T09:18:20.433329Z","amount":345735,"status":"COMPLETE","ext_ref_id":null,"payout_type":"ACH","request_data":{},"order_data":null,"product_id":"fc14e741b5004581b30e6478363414df","method":"ACH","amount_usd":345735,"amount_usd_str":"$3,457.35"}],"amount":345735,"amount_usd_str":"$3,457.35","created":"2025-08-02T09:18:20.433329Z","line_items":1,"ext_ref_id":null},{"bp_payouts":[{"uuid":"63ce1787087248978919015c8fcd5ab9","debit_account_uuid":"3a058056da85493f9b7cdfe375aad0e0","cashout_method_uuid":"602113e330cf43ae85c07d94b5100291","created":"2025-06-10T22:16:18.765668Z","amount":105778,"status":"COMPLETE","ext_ref_id":"11175997868","payout_type":"ACH","request_data":{},"order_data":null,"product_id":"fc14e741b5004581b30e6478363414df","method":"ACH","amount_usd":105778,"amount_usd_str":"$1,057.78"}],"amount":105778,"amount_usd_str":"$1,057.78","created":"2025-06-10T22:16:18.765668Z","line_items":1,"ext_ref_id":"11175997868"}]}'
+ # Doesn't validate anymore
+ # instance = Business.model_validate_json(
+ # json_data='{"id":123,"uuid":"947f6ba5250d442b9a66cde9ee33605a","name":"Example » Demo","kind":"c","tax_number":null,"contact":null,"addresses":[],"teams":[{"id":53,"uuid":"8e4197dcaefe4f1f831a02b212e6b44a","name":"Example » Demo","memberships":null,"gr_users":null,"businesses":null,"products":null}],"products":[{"id":"fc23e741b5004581b30e6478363525df","id_int":1234,"name":"Example","enabled":true,"payments_enabled":true,"created":"2025-04-14T13:25:37.279403Z","team_id":"9e4197dcaefe4f1f831a02b212e6b44a","business_id":"857f6ba6160d442b9a66cde9ee33605a","tags":[],"commission_pct":"0.050000","redirect_url":"https://pam-api-us.reppublika.com/v2/public/4970ef00-0ef7-11f0-9962-05cb6323c84c/grl/status","harmonizer_domain":"https://talk.generalresearch.com/","sources_config":{"user_defined":[{"name":"w","active":false,"banned_countries":[],"allow_mobile_ip":true,"supplier_id":null,"allow_pii_only_buyers":false,"allow_unhashed_buyers":false,"withhold_profiling":false,"pass_unconditional_eligible_unknowns":true,"address":null,"allow_vpn":null,"distribute_harmonizer_active":null}]},"session_config":{"max_session_len":600,"max_session_hard_retry":5,"min_payout":"0.14"},"payout_config":{"payout_format":null,"payout_transformation":null},"user_wallet_config":{"enabled":false,"amt":false,"supported_payout_types":["CASH_IN_MAIL","PAYPAL","TANGO"],"min_cashout":null},"user_create_config":{"min_hourly_create_limit":0,"max_hourly_create_limit":null},"offerwall_config":{},"profiling_config":{"enabled":true,"grs_enabled":true,"n_questions":null,"max_questions":10,"avg_question_count":5.0,"task_injection_freq_mult":1.0,"non_us_mult":2.0,"hidden_questions_expiration_hours":168},"user_health_config":{"banned_countries":[],"allow_ban_iphist":true},"yield_man_config":{},"balance":null,"payouts_total_str":null,"payouts_total":null,"payouts":null,"user_wallet":{"enabled":false,"amt":false,"supported_payout_types":["CASH_IN_MAIL","PAYPAL","TANGO"],"min_cashout":null}}],"bank_accounts":[],"balance":{"product_balances":[{"product_id":"fc14e741b5004581b30e6478363414df","last_event":null,"bp_payment_credit":780251,"adjustment_credit":4678,"adjustment_debit":26446,"supplier_credit":0,"supplier_debit":451513,"user_bonus_credit":0,"user_bonus_debit":0,"issued_payment":0,"payout":780251,"payout_usd_str":"$7,802.51","adjustment":-21768,"expense":0,"net":758483,"payment":451513,"payment_usd_str":"$4,515.13","balance":306970,"retainer":76742,"retainer_usd_str":"$767.42","available_balance":230228,"available_balance_usd_str":"$2,302.28","recoup":0,"recoup_usd_str":"$0.00","adjustment_percent":0.027898714644390074}],"payout":780251,"payout_usd_str":"$7,802.51","adjustment":-21768,"expense":0,"net":758483,"net_usd_str":"$7,584.83","payment":451513,"payment_usd_str":"$4,515.13","balance":306970,"balance_usd_str":"$3,069.70","retainer":76742,"retainer_usd_str":"$767.42","available_balance":230228,"available_balance_usd_str":"$2,302.28","adjustment_percent":0.027898714644390074,"recoup":0,"recoup_usd_str":"$0.00"},"payouts_total_str":"$4,515.13","payouts_total":451513,"payouts":[{"bp_payouts":[{"uuid":"40cf2c3c341e4f9d985be4bca43e6116","debit_account_uuid":"3a058056da85493f9b7cdfe375aad0e0","cashout_method_uuid":"602113e330cf43ae85c07d94b5100291","created":"2025-08-02T09:18:20.433329Z","amount":345735,"status":"COMPLETE","ext_ref_id":null,"payout_type":"ACH","request_data":{},"order_data":null,"product_id":"fc14e741b5004581b30e6478363414df","method":"ACH","amount_usd":345735,"amount_usd_str":"$3,457.35"}],"amount":345735,"amount_usd_str":"$3,457.35","created":"2025-08-02T09:18:20.433329Z","line_items":1,"ext_ref_id":null},{"bp_payouts":[{"uuid":"63ce1787087248978919015c8fcd5ab9","debit_account_uuid":"3a058056da85493f9b7cdfe375aad0e0","cashout_method_uuid":"602113e330cf43ae85c07d94b5100291","created":"2025-06-10T22:16:18.765668Z","amount":105778,"status":"COMPLETE","ext_ref_id":"11175997868","payout_type":"ACH","request_data":{},"order_data":null,"product_id":"fc14e741b5004581b30e6478363414df","method":"ACH","amount_usd":105778,"amount_usd_str":"$1,057.78"}],"amount":105778,"amount_usd_str":"$1,057.78","created":"2025-06-10T22:16:18.765668Z","line_items":1,"ext_ref_id":"11175997868"}]}'
+ # )
+ # assert isinstance(instance, Business)
+
+ # Make manually
+ b = Business(
+ id=123,
+ uuid=uuid4().hex,
+ name="Example",
+ addresses=[
+ BusinessAddress(
+ uuid=uuid4().hex,
+ city="xxx",
+ line_1="xxx",
+ state="fl",
+ business_id=123,
+ )
+ ],
+ kind=BusinessType.COMPANY,
+ teams=[Team(uuid=uuid4().hex, name="Example » Demo")],
+ products=[],
+ bank_accounts=[],
+ )
+ assert isinstance(b, Business)
+
+ ext_ref_id = uuid4().hex
+ bpe = BusinessPayoutEvent(
+ business_id=uuid4().hex,
+ amount=USDCent(100_00),
+ payout_type=PayoutType.ACH,
+ ext_ref_id=ext_ref_id,
)
+ bpe.bp_payouts = [
+ BrokerageProductPayoutEvent(
+ product_id=uuid4().hex,
+ payout_type=PayoutType.ACH,
+ amount=USDCent(47_00),
+ cashout_method_uuid=uuid4().hex,
+ debit_account_uuid=uuid4().hex,
+ ext_ref_id=ext_ref_id,
+ ),
+ BrokerageProductPayoutEvent(
+ product_id=uuid4().hex,
+ payout_type=PayoutType.ACH,
+ amount=USDCent(53_00),
+ cashout_method_uuid=uuid4().hex,
+ debit_account_uuid=uuid4().hex,
+ ext_ref_id=ext_ref_id,
+ ),
+ ]
+
+ # Test validations (amount sum)
+ with pytest.raises(
+ ValidationError,
+ match="BusinessPayoutEvent.amount must equal the sum of bp_payouts amounts",
+ ):
+ bpe.bp_payouts = [
+ BrokerageProductPayoutEvent(
+ product_id=uuid4().hex,
+ payout_type=PayoutType.ACH,
+ amount=USDCent(47_00),
+ cashout_method_uuid=uuid4().hex,
+ debit_account_uuid=uuid4().hex,
+ ext_ref_id=ext_ref_id,
+ )
+ ]
+
+ with pytest.raises(
+ ValidationError,
+ match="All BrokerageProductPayoutEvent.ext_ref_id values must equal",
+ ):
+ bpe.bp_payouts = [
+ BrokerageProductPayoutEvent(
+ product_id=uuid4().hex,
+ payout_type=PayoutType.ACH,
+ amount=USDCent(100_00),
+ cashout_method_uuid=uuid4().hex,
+ debit_account_uuid=uuid4().hex,
+ ext_ref_id="a different value",
+ )
+ ]
- assert isinstance(instance, Business)
+ with pytest.raises(
+ ValidationError, match="All BrokerageProductPayoutEvent.payout_type values"
+ ):
+ bpe.bp_payouts = [
+ BrokerageProductPayoutEvent(
+ product_id=uuid4().hex,
+ payout_type=PayoutType.PAYPAL,
+ amount=USDCent(100_00),
+ cashout_method_uuid=uuid4().hex,
+ debit_account_uuid=uuid4().hex,
+ ext_ref_id=ext_ref_id,
+ )
+ ]
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 39469dc..97abf0c 100644
--- a/tests/models/thl/test_product.py
+++ b/tests/models/thl/test_product.py
@@ -2,9 +2,10 @@ 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
from decimal import Decimal
-from typing import Callable
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
@@ -12,18 +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.mergers.pop_ledger import PopLedgerMerge
-from generalresearch.managers.thl.ledger_manager.thl_ledger import (
- ThlLedgerManager,
-)
-from generalresearch.managers.thl.product import ProductManager
-from generalresearch.models 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.product import (
- BrokerageProductPayoutEvent,
- BrokerageProductPayoutEventManager,
IntegrationMode,
PayoutConfig,
PayoutTransformation,
@@ -35,12 +27,27 @@ from generalresearch.models.thl.product import (
SupplyConfig,
SupplyPolicy,
)
-from generalresearch.models.thl.session import Session
-from generalresearch.models.thl.user import User
+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:
+class TestProduct:
def test_init(self):
# By default, just a Pydantic instance doesn't have an id_int
instance = Product.model_validate(
@@ -56,17 +63,19 @@ 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)
+ 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(
@@ -140,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",
@@ -156,6 +159,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 +295,10 @@ 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_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
@@ -583,14 +591,13 @@ class TestGlobalProductConfigFor:
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:
- return "30d"
+ return "30D"
@pytest.fixture
def duration(self) -> timedelta | None:
@@ -598,55 +605,75 @@ 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,
+ brokerage_product_payout_event_factory: Callable[
+ ..., BrokerageProductPayoutEvent
+ ],
+ thl_ledger_manager: ThlLedgerManager,
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,
+ ledger_collection: LedgerDFCollection,
pop_ledger_merge: PopLedgerMerge,
- delete_df_collection,
+ delete_df_collection: Callable[..., None],
):
delete_ledger_db()
create_main_accounts()
delete_df_collection(coll=ledger_collection)
- 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)
- brokerage_product_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
+ bp_wallet = thl_ledger_manager.get_account_or_create_bp_wallet(product=p1)
+ thl_ledger_manager.get_account_or_create_user_wallet(user=u1)
- 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,
)
@@ -656,7 +683,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,
)
@@ -670,7 +697,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
@@ -680,14 +707,21 @@ class TestProductFinancials:
# -- Now pay them out...
- bp_payout_factory(
+ from generalresearch.currency import USDCent
+
+ brokerage_product_payout_event_factory(
product=p1,
amount=USDCent(50),
created=start + timedelta(days=3),
- 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)
@@ -699,7 +733,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,
)
@@ -713,24 +747,29 @@ 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
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!.
- bp_payout_factory(
+ brokerage_product_payout_event_factory(
product=p1,
amount=USDCent(5),
created=start + timedelta(days=4),
- 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)
@@ -742,7 +781,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,
)
@@ -756,7 +795,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
@@ -766,14 +805,13 @@ class TestProductFinancials:
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:
- return "30d"
+ return "30D"
@pytest.fixture
def duration(self) -> timedelta | None:
@@ -783,18 +821,20 @@ class TestProductBalance:
self,
product: Product,
mnt_filepath: GRLDatasets,
- thl_lm: ThlLedgerManager,
+ thl_ledger_manager: ThlLedgerManager,
client_no_amm: DaskClient,
- delete_ledger_db,
- create_main_accounts,
- delete_df_collection,
- ledger_collection,
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ delete_df_collection: Callable[..., None],
+ 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,
+ brokerage_product_payout_event_factory: Callable[
+ ..., BrokerageProductPayoutEvent
+ ],
+ payout_event_manager: PayoutEventManager,
):
# Now let's load it up and actually test some things
delete_ledger_db()
@@ -813,21 +853,18 @@ 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)
- bp_payout_factory(
+ brokerage_product_payout_event_factory(
product=product,
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)
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)
@@ -835,18 +872,20 @@ class TestProductBalance:
self,
product: Product,
mnt_filepath: GRLDatasets,
- thl_lm: ThlLedgerManager,
+ thl_ledger_manager: ThlLedgerManager,
client_no_amm: DaskClient,
- delete_ledger_db,
- create_main_accounts,
- delete_df_collection,
- ledger_collection,
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ delete_df_collection: Callable[..., None],
+ ledger_collection: LedgerDFCollection,
user_factory: Callable[..., User],
- session_with_tx_factory,
+ session_with_tx_factory: Callable[..., None],
pop_ledger_merge: PopLedgerMerge,
start: datetime,
- bp_payout_factory,
- payout_event_manager,
+ brokerage_product_payout_event_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
@@ -872,31 +911,29 @@ 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)
- bp_payout_factory(
+ brokerage_product_payout_event_factory(
product=product,
amount=USDCent(71),
ext_ref_id=uuid4().hex,
- created=datetime.now(tz=timezone.utc),
- skip_wallet_balance_check=True,
- skip_one_per_day_check=True,
+ created=datetime.now(tz=UTC),
)
ledger_collection.initial_load(client=None, sync=True)
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:
-
@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:
- return "30d"
+ return "30D"
@pytest.fixture
def duration(self) -> timedelta | None:
@@ -906,14 +943,14 @@ class TestProductPOPFinancial:
self,
product: Product,
mnt_filepath: GRLDatasets,
- thl_lm: ThlLedgerManager,
+ thl_ledger_manager: ThlLedgerManager,
client_no_amm: DaskClient,
- delete_ledger_db,
- create_main_accounts,
- delete_df_collection,
- ledger_collection,
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ delete_df_collection: Callable[..., None],
+ ledger_collection: LedgerDFCollection,
user_factory: Callable[..., User],
- session_with_tx_factory,
+ session_with_tx_factory: Callable[..., None],
pop_ledger_merge: PopLedgerMerge,
start: datetime,
):
@@ -942,7 +979,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,
@@ -962,14 +999,13 @@ class TestProductPOPFinancial:
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:
- return "30d"
+ return "30D"
@pytest.fixture
def duration(self) -> timedelta | None:
@@ -978,17 +1014,17 @@ class TestProductCache:
def test_basic(
self,
product: Product,
- mnt_filepath,
- thl_lm,
+ mnt_filepath: GRLDatasets,
+ thl_ledger_manager: ThlLedgerManager,
client_no_amm: DaskClient,
- thl_redis_config,
- brokerage_product_payout_event_manager,
- delete_ledger_db,
- create_main_accounts,
- delete_df_collection,
- ledger_collection,
+ thl_redis_config: RedisConfig,
+ brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager,
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ delete_df_collection: Callable[..., None],
+ ledger_collection: LedgerDFCollection,
user_factory: Callable[..., User],
- session_with_tx_factory,
+ session_with_tx_factory: Callable[..., None],
pop_ledger_merge: PopLedgerMerge,
start: datetime,
):
@@ -1003,7 +1039,7 @@ class TestProductCache:
assert res is None
with pytest.raises(expected_exception=AssertionError):
product.set_cache(
- thl_lm=thl_lm,
+ thl_lm=thl_ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
bp_pem=brokerage_product_payout_event_manager,
@@ -1025,7 +1061,7 @@ class TestProductCache:
# Now try again with everything in place
product.set_cache(
- thl_lm=thl_lm,
+ thl_lm=thl_ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
bp_pem=brokerage_product_payout_event_manager,
@@ -1050,21 +1086,23 @@ class TestProductCache:
self,
product: Product,
mnt_filepath: GRLDatasets,
- thl_lm,
+ thl_ledger_manager: ThlLedgerManager,
client_no_amm: DaskClient,
- thl_redis_config,
- brokerage_product_payout_event_manager,
- delete_ledger_db,
- create_main_accounts,
- delete_df_collection,
- ledger_collection,
+ thl_redis_config: RedisConfig,
+ brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager,
+ delete_ledger_db: Callable[..., None],
+ create_main_accounts: Callable[..., None],
+ delete_df_collection: Callable[..., None],
+ ledger_collection: LedgerDFCollection,
user_factory: Callable[..., User],
- session_with_tx_factory,
+ session_with_tx_factory: Callable[..., None],
pop_ledger_merge: PopLedgerMerge,
start: datetime,
- bp_payout_factory,
- payout_event_manager,
- adj_to_fail_with_tx_factory,
+ brokerage_product_payout_event_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,14 +1121,11 @@ class TestProductCache:
)
# 2. Payout
- payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
- bp_payout_factory(
+ brokerage_product_payout_event_factory(
product=product,
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
@@ -1104,7 +1139,7 @@ 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,
diff --git a/tests/models/thl/test_product_userwalletconfig.py b/tests/models/thl/test_product_userwalletconfig.py
index 614fc0a..9fb9c73 100644
--- a/tests/models/thl/test_product_userwalletconfig.py
+++ b/tests/models/thl/test_product_userwalletconfig.py
@@ -1,13 +1,15 @@
+from __future__ import annotations
+
from itertools import groupby
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):
+def all_equal(iterable: list[str]) -> bool:
g = groupby(iterable)
return next(g, True) and not next(g, False)
@@ -41,13 +43,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..34902e2 100644
--- a/tests/models/thl/test_soft_pair.py
+++ b/tests/models/thl/test_soft_pair.py
@@ -1,12 +1,14 @@
-from generalresearch.models import Source
+from __future__ import annotations
+
+from generalresearch.models.definitions 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 d32875c..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,24 +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(
@@ -226,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",
@@ -266,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(
{
@@ -304,9 +284,6 @@ class TestUpkQuestionValidateAnswer:
)
def test_validate_answer_MA(self):
- from generalresearch.models.thl.profiling.upk_question import (
- UpkQuestion,
- )
question = UpkQuestion.model_validate(
{
@@ -376,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 943ae8e..68b413c 100644
--- a/tests/models/thl/test_user.py
+++ b/tests/models/thl/test_user.py
@@ -1,25 +1,35 @@
+from __future__ import annotations
+
import json
-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 rand_choice
from random import randint
+from typing import TYPE_CHECKING
from uuid import uuid4
import pytest
from pydantic import ValidationError
+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
+ from generalresearch.models.thl.userhealth import AuditLog
+
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 +54,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 +61,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 +68,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 +76,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 +86,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 +94,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 +106,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 +113,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 +135,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 +145,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 +157,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 +171,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 +182,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)
@@ -197,12 +194,11 @@ 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)
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 +206,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 +220,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,9 +228,8 @@ 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 = 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)
@@ -253,7 +246,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 +267,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 +279,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 +287,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)
@@ -310,12 +299,11 @@ 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)
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 +311,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 +328,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 +338,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 +354,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,33 +364,29 @@ 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=timezone.utc)
+ dt = datetime.now(tz=UTC)
user.created = dt
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))
+ User(user_id=self.user_id, created=datetime.now(tz=None)) # noqa
assert "1 validation error for User" in str(cm.value)
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:
- user.created = datetime.now(tz=None)
+ user.created = datetime.now(tz=None) # noqa
assert "1 validation error for User" in str(cm.value)
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,20 +397,18 @@ 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=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)
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=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)
@@ -441,33 +419,29 @@ 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=timezone.utc)
+ dt = datetime.now(tz=UTC)
user.last_seen = dt
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))
+ User(user_id=self.user_id, last_seen=datetime.now(tz=None)) # noqa
assert "1 validation error for User" in str(cm.value)
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:
- user.last_seen = datetime.now(tz=None)
+ user.last_seen = datetime.now(tz=None) # noqa
assert "1 validation error for User" in str(cm.value)
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,20 +452,18 @@ 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=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)
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=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)
@@ -502,7 +474,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 +481,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,20 +517,18 @@ class TestUserTiming:
user_id = randint(1, 2**30)
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
assert user.last_seen == last_seen
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)
@@ -572,7 +540,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 +547,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 +560,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
@@ -602,7 +567,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,
)
@@ -615,7 +580,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
@@ -623,7 +587,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,10 +597,11 @@ class TestUserSerialization:
assert not d.get("blocked")
assert d.get("product") is None
- assert d.get("created").tzinfo == timezone.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
@@ -644,41 +609,51 @@ 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 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,
+ 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
def test_transactions(
- self, user_factory, thl_lm, session_with_tx_factory, product_user_wallet_yes
+ self,
+ user_factory: Callable[..., User],
+ thl_ledger_manager: ThlLedgerManager,
+ session_with_tx_factory: Callable[..., None],
+ 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 596849c..b8a0be3 100644
--- a/tests/models/thl/test_user_iphistory.py
+++ b/tests/models/thl/test_user_iphistory.py
@@ -1,4 +1,6 @@
-from datetime import datetime, timedelta, timezone
+from __future__ import annotations
+
+from datetime import UTC, datetime, timedelta
from generalresearch.models.thl.user_iphistory import (
UserIPHistory,
@@ -8,7 +10,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_user_metadata.py b/tests/models/thl/test_user_metadata.py
index 3d851dc..7e84f3e 100644
--- a/tests/models/thl/test_user_metadata.py
+++ b/tests/models/thl/test_user_metadata.py
@@ -1,6 +1,8 @@
+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_user_streak.py b/tests/models/thl/test_user_streak.py
index 0cacd3e..8300474 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,
@@ -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 8398c81..61ca11d 100644
--- a/tests/models/thl/test_wall.py
+++ b/tests/models/thl/test_wall.py
@@ -1,11 +1,13 @@
-from datetime import datetime, timedelta, timezone
+from __future__ import annotations
+
+from datetime import UTC, datetime, timedelta
from decimal import Decimal
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,
@@ -27,8 +29,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 +47,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 +60,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 +73,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 +89,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 +106,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 +119,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 +132,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 +147,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..40d3619 100644
--- a/tests/models/thl/test_wall_session.py
+++ b/tests/models/thl/test_wall_session.py
@@ -1,9 +1,11 @@
-from datetime import datetime, timedelta, timezone
+from __future__ import annotations
+
+from datetime import UTC, datetime, timedelta
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
@@ -12,7 +14,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 +26,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 +55,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 +100,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 +121,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 +141,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 +170,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 +208,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 +231,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/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
diff --git a/tests/test_postgres.py b/tests/test_postgres.py
index 3b3ddd0..de3f5d8 100644
--- a/tests/test_postgres.py
+++ b/tests/test_postgres.py
@@ -1,18 +1,21 @@
import socket
import subprocess
-from typing import Callable
+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:
with socket.create_connection((host, port), timeout=timeout):
return True
- except (socket.timeout, ConnectionRefusedError, OSError):
+ except (TimeoutError, ConnectionRefusedError, OSError):
return False
@@ -65,4 +68,32 @@ class TestPostgresDjangoCreation:
WHERE table_schema = 'public';
""")
assert len(res) == 1
- assert res[0]["count"] == 56
+ assert res[0]["count"] == 57
+
+ def test_django_tables_only_gr(self, gr_db: PostgresConfig):
+ """
+ IMPORTANT: This can't really run with only the GR tables,
+ that's because we have most of the database init fixtures
+ as session scoped; and we can't ensure that this will
+ run before any test that depends on the core thl
+ migrations
+ """
+
+ 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"] == 65
+
+ 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"] == 65
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