aboutsummaryrefslogtreecommitdiff
path: root/tests/managers/thl/test_product_prod.py
diff options
context:
space:
mode:
Diffstat (limited to 'tests/managers/thl/test_product_prod.py')
-rw-r--r--tests/managers/thl/test_product_prod.py18
1 files changed, 13 insertions, 5 deletions
diff --git a/tests/managers/thl/test_product_prod.py b/tests/managers/thl/test_product_prod.py
index 0f622b6..8734210 100644
--- a/tests/managers/thl/test_product_prod.py
+++ b/tests/managers/thl/test_product_prod.py
@@ -1,8 +1,12 @@
+from __future__ import annotations
+
import logging
+from collections.abc import Callable
from uuid import uuid4
import pytest
+from generalresearch.managers.thl.product import ProductManager
from generalresearch.models.thl.product import Product
logger = logging.getLogger()
@@ -10,7 +14,9 @@ logger = logging.getLogger()
class TestProductManagerGetMethods:
- def test_get_by_uuid(self, product_manager: ProductManager, product_factory):
+ def test_get_by_uuid(
+ self, product_manager: ProductManager, product_factory: Callable[..., Product]
+ ):
# Just test that we load properly
for p in [product_factory(), product_factory(), product_factory()]:
instance = product_manager.get_by_uuid(product_uuid=p.id)
@@ -22,7 +28,9 @@ class TestProductManagerGetMethods:
product_manager.get_by_uuid(product_uuid=uuid4().hex)
assert "product not found" in str(cm.value)
- def test_get_by_uuids(self, product_manager: ProductManager, product_factory):
+ def test_get_by_uuids(
+ self, product_manager: ProductManager, product_factory: Callable[..., Product]
+ ):
products = [product_factory(), product_factory(), product_factory()]
cnt = len(products)
res = product_manager.get_by_uuids(product_uuids=[p.id for p in products])
@@ -43,7 +51,7 @@ class TestProductManagerGetMethods:
assert "invalid uuid passed" in str(cm.value)
def test_get_by_uuid_if_exists(
- self, product_factory: Callable[..., Product], product_manager
+ self, product_factory: Callable[..., Product], product_manager: ProductManager
):
products = [product_factory(), product_factory(), product_factory()]
@@ -54,7 +62,7 @@ class TestProductManagerGetMethods:
assert instance is None
def test_get_by_uuids_if_exists(
- self, product_manager: ProductManager, product_factory
+ self, product_manager: ProductManager, product_factory: Callable[..., Product]
):
products = [product_factory(), product_factory(), product_factory()]
@@ -78,7 +86,7 @@ class TestProductManagerGetMethods:
class TestProductManagerGetAll:
@pytest.mark.skip(reason="TODO")
- def test_get_ALL_by_ids(self, product_manager):
+ def test_get_ALL_by_ids(self, product_manager: ProductManager):
products = product_manager.get_all(rand_limit=50)
logger.info(f"Fetching {len(products)} product uuids")
# todo: once timebucks stops spamming broken accounts, fetch more