aboutsummaryrefslogtreecommitdiff
path: root/tests/models
diff options
context:
space:
mode:
Diffstat (limited to 'tests/models')
-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
51 files changed, 1814 insertions, 2047 deletions
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(