aboutsummaryrefslogtreecommitdiff
path: root/tests/models
diff options
context:
space:
mode:
authorMax Nanis2026-08-21 11:31:13 -0700
committerMax Nanis2026-08-21 11:31:13 -0700
commitff538deb290c851364df85aa6ccc64bc3719ef64 (patch)
treec04c3573ab6adab70971fef00d3420a5da5f9619 /tests/models
parent4fe0f6b5e0f0c744902e4c3ab8940e23a6f8a2e1 (diff)
parentc4a44873540ca4c0a3ab19b9beef4cfc6e0252a7 (diff)
downloadgeneralresearch-ff538deb290c851364df85aa6ccc64bc3719ef64.tar.gz
generalresearch-ff538deb290c851364df85aa6ccc64bc3719ef64.zip
Merge branch 'master' into dev
Diffstat (limited to 'tests/models')
-rw-r--r--tests/models/admin/test_report_request.py23
-rw-r--r--tests/models/custom_types/test_aware_datetime.py7
-rw-r--r--tests/models/custom_types/test_dsn.py9
-rw-r--r--tests/models/custom_types/test_uuid_str.py7
-rw-r--r--tests/models/dynata/test_eligbility.py10
-rw-r--r--tests/models/gr/test_authentication.py36
-rw-r--r--tests/models/gr/test_base.py46
-rw-r--r--tests/models/gr/test_business.py75
-rw-r--r--tests/models/thl/test_product.py53
9 files changed, 165 insertions, 101 deletions
diff --git a/tests/models/admin/test_report_request.py b/tests/models/admin/test_report_request.py
index cf4c405..a80afbe 100644
--- a/tests/models/admin/test_report_request.py
+++ b/tests/models/admin/test_report_request.py
@@ -1,4 +1,4 @@
-from datetime import timezone, datetime
+from datetime import datetime, timezone
import pandas as pd
import pytest
@@ -6,7 +6,7 @@ from pydantic import ValidationError
class TestReportRequest:
- def test_base(self, utc_60days_ago):
+ def test_base(self):
from generalresearch.models.admin.request import (
ReportRequest,
ReportType,
@@ -24,7 +24,7 @@ class TestReportRequest:
rr1 = ReportRequest.model_validate(
{
"start": datetime(
- year=datetime.now().year,
+ year=datetime.now(tz=timezone.utc).year,
month=1,
day=1,
hour=0,
@@ -43,7 +43,7 @@ class TestReportRequest:
rr2 = ReportRequest.model_validate(
{
"start": datetime(
- year=datetime.now().year,
+ year=datetime.now(tz=timezone.utc).year,
month=1,
day=1,
hour=6,
@@ -81,29 +81,30 @@ class TestReportRequest:
# interval='1d', include_open_bucket=True,
# start_floor=datetime.datetime(2025, 7, 9, 0, 0, tzinfo=datetime.timezone.utc)).start_floor
- def test_start_end_range(self, utc_90days_ago, utc_30days_ago):
+ def test_start_end_range(self, utc_90days_ago: datetime, utc_30days_ago: datetime):
from generalresearch.models.admin.request import ReportRequest
- with pytest.raises(expected_exception=ValidationError) as cm:
+ with pytest.raises(expected_exception=ValidationError):
ReportRequest.model_validate(
{"start": utc_30days_ago, "end": utc_90days_ago}
)
- with pytest.raises(expected_exception=ValidationError) as cm:
+ with pytest.raises(expected_exception=ValidationError):
ReportRequest.model_validate(
{
- "start": datetime(year=1990, month=1, day=1),
- "end": datetime(year=1950, month=1, day=1),
+ "start": datetime(year=1990, month=1, day=1, tzinfo=timezone.utc),
+ "end": datetime(year=1950, month=1, day=1, tzinfo=timezone.utc),
}
)
def test_start_end_range_tz(self):
- from generalresearch.models.admin.request import ReportRequest
from zoneinfo import ZoneInfo
+ from generalresearch.models.admin.request import ReportRequest
+
pacific_tz = ZoneInfo("America/Los_Angeles")
- with pytest.raises(expected_exception=ValidationError) as cm:
+ with pytest.raises(expected_exception=ValidationError):
ReportRequest.model_validate(
{
"start": datetime(year=2000, month=1, day=1, tzinfo=pacific_tz),
diff --git a/tests/models/custom_types/test_aware_datetime.py b/tests/models/custom_types/test_aware_datetime.py
index a23413c..530142e 100644
--- a/tests/models/custom_types/test_aware_datetime.py
+++ b/tests/models/custom_types/test_aware_datetime.py
@@ -1,10 +1,11 @@
+from __future__ import annotations
+
import logging
from datetime import datetime, timezone
-from typing import Optional
import pytest
import pytz
-from pydantic import BaseModel, ValidationError, Field
+from pydantic import BaseModel, Field, ValidationError
from generalresearch.models.custom_types import AwareDatetimeISO
@@ -12,7 +13,7 @@ logger = logging.getLogger()
class AwareDatetimeISOModel(BaseModel):
- dt_optional: Optional[AwareDatetimeISO] = Field(default=None)
+ dt_optional: AwareDatetimeISO | None = Field(default=None)
dt: AwareDatetimeISO
diff --git a/tests/models/custom_types/test_dsn.py b/tests/models/custom_types/test_dsn.py
index b37f2c4..16e1f83 100644
--- a/tests/models/custom_types/test_dsn.py
+++ b/tests/models/custom_types/test_dsn.py
@@ -2,13 +2,11 @@ from typing import Optional
from uuid import uuid4
import pytest
-from pydantic import BaseModel, ValidationError, Field
-from pydantic import MySQLDsn
+from pydantic import BaseModel, Field, MySQLDsn, ValidationError
from pydantic_core import Url
from generalresearch.models.custom_types import DaskDsn, SentryDsn
-
# --- Test Pydantic Models ---
@@ -27,7 +25,7 @@ class TestDaskDsn:
from dask.distributed import Client
m = SettingsModel(dask="tcp://dask-scheduler.internal")
-
+ assert isinstance(m.dask, Url)
assert m.dask.scheme == "tcp"
assert m.dask.host == "dask-scheduler.internal"
assert m.dask.port == 8786
@@ -72,6 +70,7 @@ class TestDaskDsn:
def test_port(self):
m = SettingsModel(dask="tcp://dask-scheduler.internal")
+ assert isinstance(m.dask, Url)
assert m.dask.port == 8786
@@ -81,6 +80,7 @@ class TestSentryDsn:
sentry=f"https://{uuid4().hex}@12345.ingest.us.sentry.io/9876543"
)
+ assert isinstance(m.sentry, Url)
assert m.sentry.scheme == "https"
assert m.sentry.host == "12345.ingest.us.sentry.io"
assert m.sentry.port == 443
@@ -109,4 +109,5 @@ class TestSentryDsn:
def test_port(self):
test_url: str = f"https://{uuid4().hex}@12345.ingest.us.sentry.io/9876543"
m = SettingsModel(sentry=test_url)
+ assert isinstance(m.sentry, Url)
assert m.sentry.port == 443
diff --git a/tests/models/custom_types/test_uuid_str.py b/tests/models/custom_types/test_uuid_str.py
index 91af9ae..02e6a8b 100644
--- a/tests/models/custom_types/test_uuid_str.py
+++ b/tests/models/custom_types/test_uuid_str.py
@@ -1,14 +1,15 @@
-from typing import Optional
+from __future__ import annotations
+
from uuid import uuid4
import pytest
-from pydantic import BaseModel, ValidationError, Field
+from pydantic import BaseModel, Field, ValidationError
from generalresearch.models.custom_types import UUIDStr
class UUIDStrModel(BaseModel):
- uuid_optional: Optional[UUIDStr] = Field(default_factory=lambda: uuid4().hex)
+ uuid_optional: UUIDStr | None = Field(default_factory=lambda: uuid4().hex)
uuid: UUIDStr
diff --git a/tests/models/dynata/test_eligbility.py b/tests/models/dynata/test_eligbility.py
index 23437f5..736c971 100644
--- a/tests/models/dynata/test_eligbility.py
+++ b/tests/models/dynata/test_eligbility.py
@@ -5,10 +5,10 @@ class TestEligibility:
def test_evaluate_task_criteria(self):
from generalresearch.models.dynata.survey import (
- DynataQuotaGroup,
DynataFilterGroup,
- DynataSurvey,
+ DynataQuotaGroup,
DynataRequirements,
+ DynataSurvey,
)
filters = [[["a", "b"], ["c", "d"]], [["e"], ["f"]]]
@@ -137,10 +137,10 @@ class TestEligibility:
def test_soft_pair(self):
from generalresearch.models.dynata.survey import (
- DynataQuotaGroup,
DynataFilterGroup,
- DynataSurvey,
+ DynataQuotaGroup,
DynataRequirements,
+ DynataSurvey,
)
filters = [[["a", "b"], ["c", "d"]], [["e"], ["f"]]]
@@ -186,7 +186,7 @@ class TestEligibility:
}
)
assert task.passes_filters(criteria_evaluation)
- passes, condition_hashes = task.passes_filters_soft(criteria_evaluation)
+ passes, _ = task.passes_filters_soft(criteria_evaluation)
assert passes
# make 'e' & 'f' None, we don't pass the 2nd filtergroup
diff --git a/tests/models/gr/test_authentication.py b/tests/models/gr/test_authentication.py
index e906d8c..6c84a5d 100644
--- a/tests/models/gr/test_authentication.py
+++ b/tests/models/gr/test_authentication.py
@@ -3,17 +3,20 @@ import json
import os
from datetime import datetime, timezone
from random import randint
+from typing import Callable
from uuid import uuid4
import pytest
+from generalresearch.models.gr.authentication import GRUser
+from generalresearch.models.gr.team import Membership, Team
+
SSO_ISSUER = ""
class TestGRUser:
- def test_init(self, gr_user):
- from generalresearch.models.gr.authentication import GRUser
+ def test_init(self, gr_user: GRUser):
assert isinstance(gr_user, GRUser)
assert not gr_user.is_superuser
@@ -26,8 +29,7 @@ class TestGRUser:
def test_businesses(self):
pass
- def test_teams(self, gr_user, membership, gr_db, gr_redis_config):
- from generalresearch.models.gr.team import Team
+ def test_teams(self, gr_user: GRUser, membership, gr_db, gr_redis_config):
assert gr_user.teams is None
@@ -40,11 +42,11 @@ class TestGRUser:
def test_prefetch_team_duplicates(
self,
gr_user_token,
- gr_user,
- membership,
+ gr_user: GRUser,
+ membership: Membership,
product_factory,
membership_factory,
- team,
+ team: Team,
thl_web_rr,
gr_redis_config,
gr_db,
@@ -61,10 +63,10 @@ class TestGRUser:
def test_products(
self,
- gr_user,
+ gr_user: GRUser,
product_factory,
- team,
- membership,
+ team: Team,
+ membership: Membership,
gr_db,
thl_web_rr,
gr_redis_config,
@@ -102,12 +104,12 @@ class TestGRUserMethods:
def test_to_redis(
self,
- gr_user,
+ gr_user: GRUser,
gr_redis,
- team,
+ team: Team,
business,
product_factory,
- membership_factory,
+ membership_factory: Callable[Membership],
):
product_factory(team=team, business=business)
membership_factory(team=team, gr_user=gr_user)
@@ -122,7 +124,7 @@ class TestGRUserMethods:
def test_set_cache(
self,
- gr_user,
+ gr_user: GRUser,
gr_user_token,
gr_redis,
gr_db,
@@ -145,7 +147,7 @@ class TestGRUserMethods:
def test_set_cache_gr_user(
self,
- gr_user,
+ gr_user: GRUser,
gr_user_token,
gr_redis,
gr_redis_config,
@@ -203,9 +205,7 @@ class TestGRUserMethods:
@pytest.mark.skip
def test_set_cache_business_uuids(
self,
- gr_user,
- membership,
- gr_user_token,
+ gr_user: GRUser,
gr_redis,
gr_db,
thl_web_rr,
diff --git a/tests/models/gr/test_base.py b/tests/models/gr/test_base.py
new file mode 100644
index 0000000..a9f01a8
--- /dev/null
+++ b/tests/models/gr/test_base.py
@@ -0,0 +1,46 @@
+import subprocess
+from pathlib import Path
+from typing import Callable
+
+import pytest
+from pydantic import PostgresDsn
+
+from generalresearch.pg_helper import PostgresConfig
+
+
+class TestGRPostgresDjangoCreation:
+
+ def test_git(self, git_key_path: Path, gr_repo: Callable[..., Path]):
+ repo_path = gr_repo()
+
+ try:
+ # Run the git command inside the target directory
+ result = subprocess.run(
+ ["git", "rev-parse", "--is-inside-work-tree"],
+ cwd=repo_path,
+ capture_output=True,
+ text=True,
+ check=True,
+ )
+ # Check if the output string is exactly "true"
+ assert result.stdout.strip() == "true"
+
+ except (subprocess.CalledProcessError, FileNotFoundError) as e:
+ pytest.fail(f"Directory is not a Git repo or Git is not installed: {e}")
+
+ def test_django_creation(
+ self,
+ django_db_factory: Callable[..., None],
+ ):
+
+ dsn = django_db_factory("gr")
+ assert isinstance(dsn, PostgresDsn)
+
+ # def test_django_tables(self, thl_web_rw: PostgresConfig):
+ # res = thl_web_rw.execute_sql_query(query="""
+ # SELECT COUNT(*)
+ # FROM information_schema.tables
+ # WHERE table_schema = 'public';
+ # """)
+ # assert len(res) == 1
+ # assert res[0]["count"] == 56
diff --git a/tests/models/gr/test_business.py b/tests/models/gr/test_business.py
index e8bd06a..7a84f23 100644
--- a/tests/models/gr/test_business.py
+++ b/tests/models/gr/test_business.py
@@ -6,6 +6,7 @@ from uuid import uuid4
import pandas as pd
import pytest
+from dask.distributed import Client as DaskClient
# noinspection PyUnresolvedReferences
from distributed.utils_test import (
@@ -14,22 +15,27 @@ 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 test_utils.incite.conftest import mnt_filepath
-from test_utils.managers.conftest import (
- business_bank_account_manager,
- lm,
- thl_lm,
-)
+from generalresearch.pg_helper import PostgresConfig
class TestBusinessBankAccount:
- def test_init(self, business, business_bank_account_manager):
+ def test_init(
+ self,
+ business: Business,
+ business_bank_account_manager: BusinessBankAccountManager,
+ ):
from generalresearch.models.gr.business import (
BusinessBankAccount,
TransferMethod,
@@ -42,7 +48,13 @@ class TestBusinessBankAccount:
)
assert isinstance(instance, BusinessBankAccount)
- def test_business(self, business_bank_account, business, gr_db, gr_redis_config):
+ def test_business(
+ self,
+ business_bank_account: BusinessBankAccount,
+ business: Business,
+ gr_db,
+ gr_redis_config,
+ ):
from generalresearch.models.gr.business import Business
assert business_bank_account.business is None
@@ -56,16 +68,13 @@ class TestBusinessBankAccount:
class TestBusinessAddress:
- def test_init(self, business_address):
- from generalresearch.models.gr.business import BusinessAddress
-
+ def test_init(self, business_address: BusinessAddress):
assert isinstance(business_address, BusinessAddress)
class TestBusinessContact:
def test_init(self):
- from generalresearch.models.gr.business import BusinessContact
bc = BusinessContact(name="abc", email="test@abc.com")
assert isinstance(bc, BusinessContact)
@@ -104,7 +113,7 @@ class TestBusiness:
user_factory,
session_with_tx_factory,
pop_ledger_merge,
- client_no_amm,
+ client_no_amm: DaskClient,
ledger_collection,
mnt_filepath,
create_main_accounts,
@@ -220,11 +229,11 @@ class TestBusiness:
def test_balance(
self,
- business,
+ business: Business,
mnt_filepath,
- client_no_amm,
- thl_web_rr,
- lm,
+ client_no_amm: DaskClient,
+ thl_web_rr: PostgresConfig,
+ ledger_manager,
pop_ledger_merge,
):
assert business.balance is None
@@ -232,7 +241,7 @@ class TestBusiness:
with pytest.raises(expected_exception=AssertionError) as cm:
business.prebuild_balance(
thl_pg_config=thl_web_rr,
- lm=lm,
+ lm=ledger_manager,
ds=mnt_filepath,
client=client_no_amm,
pop_ledger=pop_ledger_merge,
@@ -248,7 +257,7 @@ class TestBusiness:
business,
product_factory,
thl_web_rr,
- thl_lm,
+ thl_ledger_manager,
business_payout_event_manager,
):
assert business.payouts is None
@@ -256,17 +265,17 @@ class TestBusiness:
with pytest.raises(expected_exception=AssertionError) as cm:
business.prebuild_payouts(
thl_pg_config=thl_web_rr,
- thl_lm=thl_lm,
+ thl_lm=thl_ledger_manager,
bpem=business_payout_event_manager,
)
assert "Must provide product_uuids" in str(cm.value)
p = product_factory(business=business)
- thl_lm.get_account_or_create_bp_wallet(product=p)
+ thl_ledger_manager.get_account_or_create_bp_wallet(product=p)
business.prebuild_payouts(
thl_pg_config=thl_web_rr,
- thl_lm=thl_lm,
+ thl_lm=thl_ledger_manager,
bpem=business_payout_event_manager,
)
assert isinstance(business.payouts, list)
@@ -274,17 +283,17 @@ class TestBusiness:
def test_payouts(
self,
- business,
- product_factory,
+ business: Business,
+ product_factory: Callable[Product],
bp_payout_factory,
- thl_lm,
+ thl_ledger_manager,
thl_web_rr,
business_payout_event_manager,
create_main_accounts,
):
create_main_accounts()
p = product_factory(business=business)
- thl_lm.get_account_or_create_bp_wallet(product=p)
+ 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(
@@ -293,7 +302,7 @@ class TestBusiness:
business.prebuild_payouts(
thl_pg_config=thl_web_rr,
- thl_lm=thl_lm,
+ thl_lm=thl_ledger_manager,
bpem=business_payout_event_manager,
)
assert len(business.payouts) == 1
@@ -306,10 +315,12 @@ class TestBusiness:
skip_wallet_balance_check=True,
skip_one_per_day_check=True,
)
- business_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm)
+ 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_lm,
+ thl_lm=thl_ledger_manager,
bpem=business_payout_event_manager,
)
assert len(business.payouts) == 1
@@ -370,7 +381,7 @@ class TestBusiness:
self,
business,
thl_web_rr,
- thl_lm,
+ thl_ledger_manager,
mnt_filepath,
client_no_amm,
pop_ledger_merge,
@@ -496,7 +507,7 @@ class TestBusinessBalance:
mnt_filepath,
bp_payout_factory,
thl_lm,
- lm,
+ ledger_manager,
duration,
offset,
start,
diff --git a/tests/models/thl/test_product.py b/tests/models/thl/test_product.py
index 5e9b249..39469dc 100644
--- a/tests/models/thl/test_product.py
+++ b/tests/models/thl/test_product.py
@@ -1,8 +1,10 @@
+from __future__ import annotations
+
import os
import shutil
from datetime import datetime, timedelta, timezone
from decimal import Decimal
-from typing import Callable, Optional
+from typing import Callable
from uuid import uuid4
import pytest
@@ -10,7 +12,7 @@ from dask.distributed import Client as DaskClient
from pydantic import ValidationError
from generalresearch.currency import USDCent
-from generalresearch.incite import GRLDatasets
+from generalresearch.incite.base import GRLDatasets
from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge
from generalresearch.managers.thl.ledger_manager.thl_ledger import (
ThlLedgerManager,
@@ -25,6 +27,7 @@ from generalresearch.models.thl.product import (
IntegrationMode,
PayoutConfig,
PayoutTransformation,
+ PayoutTransformationPercentArgs,
Product,
ProfilingConfig,
SourceConfig,
@@ -137,6 +140,12 @@ 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",
@@ -576,7 +585,7 @@ class TestGlobalProductConfigFor:
class TestProductFinancials:
@pytest.fixture
- def start(self) -> "datetime":
+ def start(self) -> datetime:
return datetime(year=2018, month=3, day=14, hour=0, tzinfo=timezone.utc)
@pytest.fixture
@@ -584,7 +593,7 @@ class TestProductFinancials:
return "30d"
@pytest.fixture
- def duration(self) -> Optional["timedelta"]:
+ def duration(self) -> timedelta | None:
return None
def test_balance(
@@ -759,7 +768,7 @@ class TestProductFinancials:
class TestProductBalance:
@pytest.fixture
- def start(self) -> "datetime":
+ def start(self) -> datetime:
return datetime(year=2018, month=3, day=14, hour=0, tzinfo=timezone.utc)
@pytest.fixture
@@ -767,7 +776,7 @@ class TestProductBalance:
return "30d"
@pytest.fixture
- def duration(self) -> Optional["timedelta"]:
+ def duration(self) -> timedelta | None:
return None
def test_inconsistent(
@@ -783,7 +792,7 @@ class TestProductBalance:
user_factory: Callable[..., User],
session_with_tx_factory: Callable[..., Session],
pop_ledger_merge,
- start,
+ start: datetime,
bp_payout_factory,
payout_event_manager,
):
@@ -792,8 +801,6 @@ class TestProductBalance:
create_main_accounts()
delete_df_collection(coll=ledger_collection)
- from generalresearch.models.thl.user import User
-
u1: User = user_factory(product=product)
# 1. Complete and Build Parquets 1st time
@@ -827,7 +834,7 @@ class TestProductBalance:
def test_not_inconsistent(
self,
product: Product,
- mnt_filepath,
+ mnt_filepath: GRLDatasets,
thl_lm: ThlLedgerManager,
client_no_amm: DaskClient,
delete_ledger_db,
@@ -852,8 +859,6 @@ class TestProductBalance:
create_main_accounts()
delete_df_collection(coll=ledger_collection)
- from generalresearch.models.thl.user import User
-
u1: User = user_factory(product=product)
# 1. Complete and Build Parquets 1st time
@@ -886,7 +891,7 @@ class TestProductBalance:
class TestProductPOPFinancial:
@pytest.fixture
- def start(self) -> "datetime":
+ def start(self) -> datetime:
return datetime(year=2018, month=3, day=14, hour=0, tzinfo=timezone.utc)
@pytest.fixture
@@ -894,15 +899,15 @@ class TestProductPOPFinancial:
return "30d"
@pytest.fixture
- def duration(self) -> Optional["timedelta"]:
+ def duration(self) -> timedelta | None:
return None
def test_base(
self,
- product,
- mnt_filepath,
+ product: Product,
+ mnt_filepath: GRLDatasets,
thl_lm: ThlLedgerManager,
- client_no_amm,
+ client_no_amm: DaskClient,
delete_ledger_db,
create_main_accounts,
delete_df_collection,
@@ -923,8 +928,6 @@ class TestProductPOPFinancial:
create_main_accounts()
delete_df_collection(coll=ledger_collection)
- from generalresearch.models.thl.user import User
-
u1: User = user_factory(product=product)
# 1. Complete and Build Parquets 1st time
@@ -961,7 +964,7 @@ class TestProductPOPFinancial:
class TestProductCache:
@pytest.fixture
- def start(self) -> "datetime":
+ def start(self) -> datetime:
return datetime(year=2018, month=3, day=14, hour=0, tzinfo=timezone.utc)
@pytest.fixture
@@ -969,7 +972,7 @@ class TestProductCache:
return "30d"
@pytest.fixture
- def duration(self) -> Optional["timedelta"]:
+ def duration(self) -> timedelta | None:
return None
def test_basic(
@@ -1008,7 +1011,6 @@ class TestProductCache:
)
from generalresearch.models.thl.product import Product
- from generalresearch.models.thl.user import User
u1: User = user_factory(product=product)
@@ -1047,7 +1049,7 @@ class TestProductCache:
def test_neg_balance_cache(
self,
product: Product,
- mnt_filepath,
+ mnt_filepath: GRLDatasets,
thl_lm,
client_no_amm: DaskClient,
thl_redis_config,
@@ -1070,7 +1072,6 @@ class TestProductCache:
delete_df_collection(coll=ledger_collection)
from generalresearch.models.thl.product import Product
- from generalresearch.models.thl.user import User
u1: User = user_factory(product=product)
@@ -1112,10 +1113,12 @@ class TestProductCache:
# Fetch from cache and assert the instance loaded from redis
rc = thl_redis_config.create_redis_client()
- res: Optional[str] = rc.get(product.cache_key)
+ res: str | None = rc.get(product.cache_key)
assert isinstance(res, str)
p1: Product = Product.model_validate_json(res)
+ assert p1.balance
+
assert p1.balance.product_id == product.uuid
assert p1.balance.payout_usd_str == "$0.71"
assert p1.balance.adjustment == -71