aboutsummaryrefslogtreecommitdiff
path: root/tests/managers/thl/test_product.py
diff options
context:
space:
mode:
Diffstat (limited to 'tests/managers/thl/test_product.py')
-rw-r--r--tests/managers/thl/test_product.py45
1 files changed, 26 insertions, 19 deletions
diff --git a/tests/managers/thl/test_product.py b/tests/managers/thl/test_product.py
index 31b7b73..f93ac36 100644
--- a/tests/managers/thl/test_product.py
+++ b/tests/managers/thl/test_product.py
@@ -1,10 +1,15 @@
+from __future__ import annotations
+
+from collections.abc import Callable
from uuid import uuid4
import pytest
+from generalresearch.managers.thl.product import ProductManager
from generalresearch.models import Source
+from generalresearch.models.gr.team import Team
from generalresearch.models.thl.product import (
- product: Product,
+ Product,
ProfilingConfig,
SourceConfig,
SourcesConfig,
@@ -16,7 +21,7 @@ from generalresearch.models.thl.product import (
class TestProductManagerGetMethods:
- def test_get_by_uuid(self, product_manager):
+ def test_get_by_uuid(self, product_manager: ProductManager):
product: Product = product_manager.create_dummy(
product_id=uuid4().hex,
team_id=uuid4().hex,
@@ -36,10 +41,10 @@ class TestProductManagerGetMethods:
product_manager.get_by_uuid(product_uuid=uuid4().hex)
assert "product not found" in str(cm.value)
- def test_get_by_uuids(self, product_manager):
+ def test_get_by_uuids(self, product_manager: ProductManager):
cnt = 5
- product_uuids = [uuid4().hex for idx in range(cnt)]
+ product_uuids = [uuid4().hex for _ in range(cnt)]
for product_id in product_uuids:
product_manager.create_dummy(
product_id=product_id,
@@ -61,7 +66,7 @@ class TestProductManagerGetMethods:
product_manager.get_by_uuids(product_uuids=product_uuids + ["abc123"])
assert "invalid uuid" in str(cm.value)
- def test_get_by_uuid_if_exists(self, product_manager):
+ def test_get_by_uuid_if_exists(self, product_manager: ProductManager):
product: Product = product_manager.create_dummy(
product_id=uuid4().hex,
team_id=uuid4().hex,
@@ -73,7 +78,7 @@ class TestProductManagerGetMethods:
instance = product_manager.get_by_uuid_if_exists(product_uuid="abc123")
assert instance == None
- def test_get_by_uuids_if_exists(self, product_manager):
+ def test_get_by_uuids_if_exists(self, product_manager: ProductManager):
product_uuids = [uuid4().hex for _ in range(2)]
for product_id in product_uuids:
product_manager.create_dummy(
@@ -105,8 +110,8 @@ class TestProductManagerGetMethods:
# for instance in res:
# assert isinstance(instance, Product)
- def test_get_by_business_ids(self, product_manager):
- business_ids = [uuid4().hex for i in range(5)]
+ def test_get_by_business_ids(self, product_manager: ProductManager):
+ business_ids = [uuid4().hex for _ in range(5)]
product_manager.fetch_uuids(business_uuids=business_ids)
@@ -123,7 +128,7 @@ class TestProductManagerGetMethods:
class TestProductManagerCreation:
- def test_base(self, product_manager):
+ def test_base(self, product_manager: ProductManager):
instance = product_manager.create_dummy(
product_id=uuid4().hex,
team_id=uuid4().hex,
@@ -135,7 +140,7 @@ class TestProductManagerCreation:
class TestProductManagerCreate:
- def test_create_simple(self, product_manager):
+ def test_create_simple(self, product_manager: ProductManager):
# Always required: product_id, team_id, name, redirect_url
# Required internally - if not passed use default: harmonizer_domain,
# commission_pct, sources
@@ -181,9 +186,9 @@ class TestProductManager:
def test_get_by_uuid1(
self,
product_manager: ProductManager,
- team,
+ team: Team,
product: Product,
- product_factory,
+ product_factory: Callable[..., Product],
):
p1 = product_factory(team=team)
instance = product_manager.get_by_uuid(product_uuid=p1.uuid)
@@ -209,7 +214,9 @@ class TestProductManager:
assert 0 == instance.user_create_config.min_hourly_create_limit
assert instance.user_create_config.max_hourly_create_limit is None
- def test_get_by_uuid3(self, product_manager: ProductManager, product_factory):
+ def test_get_by_uuid3(
+ self, product_manager: ProductManager, product_factory: Callable[..., Product]
+ ):
p3 = product_factory()
instance = product_manager.get_by_uuid(p3.id)
assert instance.id == p3.id
@@ -225,7 +232,7 @@ class TestProductManager:
assert instance.user_create_config.max_hourly_create_limit is None
assert not instance.user_wallet_config.enabled
- def test_sources(self, product_manager):
+ def test_sources(self, product_manager: ProductManager):
user_defined = [SourceConfig(name=Source.DYNATA, active=False)]
sources_config = SourcesConfig(user_defined=user_defined)
p = product_manager.create_dummy(sources_config=sources_config)
@@ -240,7 +247,7 @@ class TestProductManager:
assert not dynata.active
assert all(x.active is True for x in p2.sources if x.name != Source.DYNATA)
- def test_global_sources(self, product_manager):
+ def test_global_sources(self, product_manager: ProductManager):
sources_config = SupplyConfig(
policies=[
SupplyPolicy(
@@ -267,7 +274,7 @@ class TestProductManager:
p2 = product_manager.get_by_uuid(p1.id)
assert p1 == p2
- def test_user_health_config(self, product_manager):
+ def test_user_health_config(self, product_manager: ProductManager):
p = product_manager.create_dummy(
user_health_config=UserHealthConfig(banned_countries=["ng", "in"])
)
@@ -278,7 +285,7 @@ class TestProductManager:
assert p2.user_health_config.banned_countries == ["in", "ng"]
assert p2.user_health_config.allow_ban_iphist
- def test_profiling_config(self, product_manager):
+ def test_profiling_config(self, product_manager: ProductManager):
p = product_manager.create_dummy(
profiling_config=ProfilingConfig(max_questions=1)
)
@@ -325,7 +332,7 @@ class TestProductManager:
class TestProductManagerUpdate:
- def test_update(self, product_manager):
+ def test_update(self, product_manager: ProductManager):
p = product_manager.create_dummy()
p.name = "new name"
p.enabled = False
@@ -346,7 +353,7 @@ class TestProductManagerUpdate:
class TestProductManagerCacheClear:
- def test_cache_clear(self, product_manager):
+ def test_cache_clear(self, product_manager: ProductManager):
p = product_manager.create_dummy()
product_manager.get_by_uuid(product_uuid=p.id)
product_manager.get_by_uuid(product_uuid=p.id)