aboutsummaryrefslogtreecommitdiff
path: root/tests/managers/thl/test_product.py
diff options
context:
space:
mode:
authorMax Nanis2026-09-02 23:31:52 -0700
committerMax Nanis2026-09-02 23:31:52 -0700
commit17ff15c06655717627da820417337c6b0b97de42 (patch)
tree0cea68d654991b3742ba7447626c5e9ee639285a /tests/managers/thl/test_product.py
parentd36994dd21a2bc025188a1ab58334915221f22cc (diff)
downloadgeneralresearch-17ff15c06655717627da820417337c6b0b97de42.tar.gz
generalresearch-17ff15c06655717627da820417337c6b0b97de42.zip
Lots more tests/managers/thl - doing all the factory organization from create_dummy
Diffstat (limited to 'tests/managers/thl/test_product.py')
-rw-r--r--tests/managers/thl/test_product.py76
1 files changed, 50 insertions, 26 deletions
diff --git a/tests/managers/thl/test_product.py b/tests/managers/thl/test_product.py
index 644dc90..81e0122 100644
--- a/tests/managers/thl/test_product.py
+++ b/tests/managers/thl/test_product.py
@@ -24,8 +24,12 @@ if TYPE_CHECKING:
class TestProductManagerGetMethods:
- def test_get_by_uuid(self, product_manager: ProductManager):
- product: Product = product_manager.create_dummy(
+ def test_get_by_uuid(
+ self,
+ product_manager: ProductManager,
+ product_factory: Callable[..., Product],
+ ):
+ product: Product = product_factory(
product_id=uuid4().hex,
team_id=uuid4().hex,
name=f"Test Product ID #{uuid4().hex[:6]}",
@@ -44,12 +48,14 @@ class TestProductManagerGetMethods:
product_manager.get_by_uuid(product_uuid=uuid4().hex)
assert "product not found" in str(cm.value)
- def test_get_by_uuids(self, product_manager: ProductManager):
+ def test_get_by_uuids(
+ self, product_factory: Callable[..., Product], product_manager: ProductManager
+ ):
cnt = 5
product_uuids = [uuid4().hex for _ in range(cnt)]
for product_id in product_uuids:
- product_manager.create_dummy(
+ product_factory(
product_id=product_id,
team_id=uuid4().hex,
name=f"Test Product ID #{uuid4().hex[:6]}",
@@ -69,8 +75,10 @@ class TestProductManagerGetMethods:
product_manager.get_by_uuids(product_uuids=product_uuids + ["abc123"])
assert "invalid uuid" in str(cm.value)
- def test_get_by_uuid_if_exists(self, product_manager: ProductManager):
- product: Product = product_manager.create_dummy(
+ def test_get_by_uuid_if_exists(
+ self, product_factory: Callable[..., Product], product_manager: ProductManager
+ ):
+ product: Product = product_factory(
product_id=uuid4().hex,
team_id=uuid4().hex,
name=f"Test Product ID #{uuid4().hex[:6]}",
@@ -81,10 +89,12 @@ class TestProductManagerGetMethods:
instance = product_manager.get_by_uuid_if_exists(product_uuid="abc123")
assert instance == None
- def test_get_by_uuids_if_exists(self, product_manager: ProductManager):
+ def test_get_by_uuids_if_exists(
+ self, product_factory: Callable[..., Product], product_manager: ProductManager
+ ):
product_uuids = [uuid4().hex for _ in range(2)]
for product_id in product_uuids:
- product_manager.create_dummy(
+ product_factory(
product_id=product_id,
team_id=uuid4().hex,
name=f"Test Product ID #{uuid4().hex[:6]}",
@@ -113,13 +123,15 @@ class TestProductManagerGetMethods:
# for instance in res:
# assert isinstance(instance, Product)
- def test_get_by_business_ids(self, product_manager: ProductManager):
+ def test_get_by_business_ids(
+ self, product_factory: Callable[..., Product], product_manager: ProductManager
+ ):
business_ids = [uuid4().hex for _ in range(5)]
product_manager.fetch_uuids(business_uuids=business_ids)
for business_id in business_ids:
- product_manager.create(
+ product_factory(
product_id=uuid4().hex,
team_id=None,
business_id=business_id,
@@ -131,8 +143,10 @@ class TestProductManagerGetMethods:
class TestProductManagerCreation:
- def test_base(self, product_manager: ProductManager):
- instance = product_manager.create_dummy(
+ def test_base(
+ self, product_factory: Callable[..., Product], product_manager: ProductManager
+ ):
+ instance = product_factory(
product_id=uuid4().hex,
team_id=uuid4().hex,
name=f"New Test Product {uuid4().hex[:6]}",
@@ -235,10 +249,12 @@ class TestProductManager:
assert instance.user_create_config.max_hourly_create_limit is None
assert not instance.user_wallet_config.enabled
- def test_sources(self, product_manager: ProductManager):
+ def test_sources(
+ self, product_factory: Callable[..., Product], product_manager: ProductManager
+ ):
user_defined = [SourceConfig(name=Source.DYNATA, active=False)]
sources_config = SourcesConfig(user_defined=user_defined)
- p = product_manager.create_dummy(sources_config=sources_config)
+ p = product_factory(sources_config=sources_config)
p2 = product_manager.get_by_uuid(p.id)
@@ -250,7 +266,9 @@ class TestProductManager:
assert not dynata.active
assert all(x.active is True for x in p2.sources if x.name != Source.DYNATA)
- def test_global_sources(self, product_manager: ProductManager):
+ def test_global_sources(
+ self, product_factory: Callable[..., Product], product_manager: ProductManager
+ ):
sources_config = SupplyConfig(
policies=[
SupplyPolicy(
@@ -261,7 +279,7 @@ class TestProductManager:
)
]
)
- p1 = product_manager.create_dummy(sources_config=sources_config)
+ p1 = product_factory(sources_config=sources_config)
p2 = product_manager.get_by_uuid(p1.id)
assert p1 == p2
@@ -277,8 +295,10 @@ class TestProductManager:
p2 = product_manager.get_by_uuid(p1.id)
assert p1 == p2
- def test_user_health_config(self, product_manager: ProductManager):
- p = product_manager.create_dummy(
+ def test_user_health_config(
+ self, product_factory: Callable[..., Product], product_manager: ProductManager
+ ):
+ p = product_factory(
user_health_config=UserHealthConfig(banned_countries=["ng", "in"])
)
@@ -288,10 +308,10 @@ class TestProductManager:
assert p2.user_health_config.banned_countries == ["in", "ng"]
assert p2.user_health_config.allow_ban_iphist
- def test_profiling_config(self, product_manager: ProductManager):
- p = product_manager.create_dummy(
- profiling_config=ProfilingConfig(max_questions=1)
- )
+ def test_profiling_config(
+ self, product_factory: Callable[..., Product], product_manager: ProductManager
+ ):
+ p = product_factory(profiling_config=ProfilingConfig(max_questions=1))
p2 = product_manager.get_by_uuid(p.id)
assert p == p2
@@ -335,8 +355,10 @@ class TestProductManager:
class TestProductManagerUpdate:
- def test_update(self, product_manager: ProductManager):
- p = product_manager.create_dummy()
+ def test_update(
+ self, product_factory: Callable[..., Product], product_manager: ProductManager
+ ):
+ p = product_factory()
p.name = "new name"
p.enabled = False
p.user_create_config = UserCreateConfig(min_hourly_create_limit=200)
@@ -356,8 +378,10 @@ class TestProductManagerUpdate:
class TestProductManagerCacheClear:
- def test_cache_clear(self, product_manager: ProductManager):
- p = product_manager.create_dummy()
+ def test_cache_clear(
+ self, product_factory: Callable[..., Product], product_manager: ProductManager
+ ):
+ p = product_factory()
product_manager.get_by_uuid(product_uuid=p.id)
product_manager.get_by_uuid(product_uuid=p.id)
product_manager.pg_config.execute_write(