diff options
| author | Max Nanis | 2026-09-02 23:31:52 -0700 |
|---|---|---|
| committer | Max Nanis | 2026-09-02 23:31:52 -0700 |
| commit | 17ff15c06655717627da820417337c6b0b97de42 (patch) | |
| tree | 0cea68d654991b3742ba7447626c5e9ee639285a | |
| parent | d36994dd21a2bc025188a1ab58334915221f22cc (diff) | |
| download | generalresearch-17ff15c06655717627da820417337c6b0b97de42.tar.gz generalresearch-17ff15c06655717627da820417337c6b0b97de42.zip | |
Lots more tests/managers/thl - doing all the factory organization from create_dummy
25 files changed, 624 insertions, 360 deletions
diff --git a/generalresearch/models/gr/team.py b/generalresearch/models/gr/team.py index aa62c5a..e10dba5 100644 --- a/generalresearch/models/gr/team.py +++ b/generalresearch/models/gr/team.py @@ -132,8 +132,8 @@ class Team(BaseModel): def prefetch_gr_users(self, gr_user_manager: GRUserManager) -> None: self.gr_users = gr_user_manager.get_by_team(team_id=self.id) - def prefetch_businesses(self, business_manager: BusinessManager) -> None: - self.businesses = business_manager.get_by_team(team_id=self.id) + def prefetch_businesses(self, gr_business_manager: BusinessManager) -> None: + self.businesses = gr_business_manager.get_by_team(team_id=self.id) def prefetch_products(self, product_manager: ProductManager) -> None: self.products = product_manager.fetch_uuids(team_uuids=[self.uuid]) @@ -273,7 +273,7 @@ class Team(BaseModel): ) -> None: self.prefetch_products(product_manager=product_manager) self.prefetch_gr_users(gr_user_manager=gr_user_manager) - self.prefetch_businesses(business_manager=gr_business_manager) + self.prefetch_businesses(gr_business_manager=gr_business_manager) self.prefetch_memberships(membership_manager=gr_membership_manager) rc = redis_config.create_redis_client() diff --git a/generalresearch/thl_django/app/test_settings.py b/generalresearch/thl_django/app/test_settings.py index f3d23af..d6ab124 100644 --- a/generalresearch/thl_django/app/test_settings.py +++ b/generalresearch/thl_django/app/test_settings.py @@ -1,7 +1,7 @@ DATABASES = { "default": { "ENGINE": "django.db.backends.postgresql", - "NAME": 'unittest-2026-09-03-728bcf', + "NAME": 'unittest-2026-09-03-44c0b4', "USER": 'jenkins', "PASSWORD": '123456789', "HOST": 'unittest-postgresql.fmt2.grl.internal', diff --git a/test_utils/grliq/conftest.py b/test_utils/grliq/conftest.py index bb1a167..f399a99 100644 --- a/test_utils/grliq/conftest.py +++ b/test_utils/grliq/conftest.py @@ -27,7 +27,8 @@ if TYPE_CHECKING: GrlIqEventManager, ) -# === Miscellaneous === + +# --- Assets --- @pytest.fixture(scope="function") @@ -48,18 +49,98 @@ def grliq_db(postgres_instance: PostgresDsn) -> PostgresConfig: ) -# === Managers === +# --- GRLIQ Data --- @pytest.fixture(scope="session") -def grliq_dm(grliq_db: PostgresConfig) -> GrlIqDataManager: +def grliq_data_manager(grliq_db: PostgresConfig) -> GrlIqDataManager: assert grliq_db.dsn.path assert "/unittest-" in grliq_db.dsn.path return GrlIqDataManager(postgres_config=grliq_db) @pytest.fixture(scope="session") -def grliq_em(grliq_db: PostgresConfig) -> GrlIqEventManager: +def grliq_dm(grliq_data_manager: GrlIqDataManager) -> GrlIqDataManager: + return grliq_data_manager + + +@pytest.fixture +def grliq_data_factory( + grliq_data_manager: GrlIqDataManager, grliq_data_list: list[dict[str, Any]] +) -> Callable[..., GrlIqData]: + + def _inner( + save: bool = True, + is_attempt_allowed: bool = True, + product_id: str | None = None, + product_user_id: str | None = None, + uuid: str | None = None, + mid: str | None = None, + created_at: datetime | None = None, + ) -> GrlIqData: + """ + Creates a dummy record in the db with a GrlIqData (data), GrlIqCheckerResults (result_data), + and GrlIqForensicCategoryResult (category_results) + :param is_attempt_allowed: Whether the attempt is allowed. + :param product_id: product_id of user + :param product_user_id: product_user_id of user + :param uuid: uuid for the grliq data record + :param mid: the thl_session:uuid / mid for the attempt. + :return: + """ + + if save: + res: GrlIqData = grliq_data_list[int(is_attempt_allowed)]["data"] + + product_id = product_id or uuid4().hex + product_user_id = product_user_id or uuid4().hex + uuid = uuid or uuid4().hex + mid = mid or uuid4().hex + created_at = created_at or datetime.now(tz=UTC) + + res["data"].product_id = product_id + res["data"].product_user_id = product_user_id + res["data"].uuid = uuid + res["data"].mid = mid + res["data"].created_at = created_at + res["result_data"].uuid = uuid + res["category_result"].uuid = uuid + + return grliq_data_manager.create( + iq_data=res["data"], + result_data=res["result_data"], + category_result=res["category_result"], + fraud_score=res["category_result"].fraud_score, + is_attempt_allowed=res["category_result"].is_attempt_allowed(), + ) + else: + raise ValueError("Unsaved GRLIQ Data not supported yet") + + return _inner + + +@pytest.fixture(scope="function") +def grliq_data(grliq_data_list: list[dict[str, Any]]) -> GrlIqData: + + g: GrlIqData = grliq_data_list[1]["data"] + + g.id = None + g.uuid = uuid4().hex + g.created_at = datetime.now(tz=UTC) + g.timestamp = g.created_at - timedelta(seconds=10) + return g + + +@pytest.fixture(scope="function") +def unsaved_grliq_data(grliq_data_list: list[dict[str, Any]]) -> GrlIqData: + raise ValueError("Not supported") + + +# --- GRLIQ Event --- + + +@pytest.fixture(scope="session") +def grliq_event_manager(grliq_db: PostgresConfig) -> GrlIqEventManager: assert grliq_db.dsn.path assert "/unittest-" in grliq_db.dsn.path @@ -71,16 +152,36 @@ def grliq_em(grliq_db: PostgresConfig) -> GrlIqEventManager: @pytest.fixture(scope="session") -def grliq_crr(grliq_db: PostgresConfig) -> GrlIqCategoryResultsReader: +def grliq_em(grliq_event_manager: GrlIqEventManager) -> GrlIqEventManager: + return grliq_event_manager + + +# --- GRLIQ Category Results Reader --- + + +@pytest.fixture(scope="session") +def grliq_category_results_reader( + grliq_db: PostgresConfig, +) -> GrlIqCategoryResultsReader: assert grliq_db.dsn.path assert "/unittest-" in grliq_db.dsn.path return GrlIqCategoryResultsReader(postgres_config=grliq_db) +@pytest.fixture(scope="session") +def grliq_crr( + grliq_category_results_reader: GrlIqCategoryResultsReader, +) -> GrlIqCategoryResultsReader: + return grliq_category_results_reader + + # === Models === +# === Miscellaneous === + + @pytest.fixture(scope="session") def grliq_data_list() -> list[dict[str, Any]]: return [ @@ -111,66 +212,3 @@ def grliq_data_list() -> list[dict[str, Any]]: "is_attempt_allowed": True, }, ] - - -@pytest.fixture(scope="function") -def grliq_data(grliq_data_list: list[dict[str, Any]]) -> GrlIqData: - - g: GrlIqData = grliq_data_list[1]["data"] - - g.id = None - g.uuid = uuid4().hex - g.created_at = datetime.now(tz=UTC) - g.timestamp = g.created_at - timedelta(seconds=10) - return g - - -@pytest.fixture -def grliq_data_factory( - grliq_dm: GrlIqDataManager, grliq_data_list: list[dict[str, Any]] -) -> Callable[..., GrlIqData]: - - def _inner( - is_attempt_allowed: bool = True, - product_id: str | None = None, - product_user_id: str | None = None, - uuid: str | None = None, - mid: str | None = None, - created_at: datetime | None = None, - ) -> GrlIqData: - """ - Creates a dummy record in the db with a GrlIqData (data), GrlIqCheckerResults (result_data), - and GrlIqForensicCategoryResult (category_results) - :param is_attempt_allowed: Whether the attempt is allowed. - :param product_id: product_id of user - :param product_user_id: product_user_id of user - :param uuid: uuid for the grliq data record - :param mid: the thl_session:uuid / mid for the attempt. - :return: - """ - - res: GrlIqData = grliq_data_list[int(is_attempt_allowed)]["data"] - - product_id = product_id or uuid4().hex - product_user_id = product_user_id or uuid4().hex - uuid = uuid or uuid4().hex - mid = mid or uuid4().hex - created_at = created_at or datetime.now(tz=UTC) - - res["data"].product_id = product_id - res["data"].product_user_id = product_user_id - res["data"].uuid = uuid - res["data"].mid = mid - res["data"].created_at = created_at - res["result_data"].uuid = uuid - res["category_result"].uuid = uuid - - return grliq_dm.create( - iq_data=res["data"], - result_data=res["result_data"], - category_result=res["category_result"], - fraud_score=res["category_result"].fraud_score, - is_attempt_allowed=res["category_result"].is_attempt_allowed(), - ) - - return _inner diff --git a/test_utils/managers/conftest.py b/test_utils/managers/conftest.py index ff088c2..3e7b304 100644 --- a/test_utils/managers/conftest.py +++ b/test_utils/managers/conftest.py @@ -56,16 +56,6 @@ def ip_geoname_manager(thl_web_rw: PostgresConfig) -> IPGeonameManager: @pytest.fixture(scope="session") -def ip_information_manager(thl_web_rw: PostgresConfig) -> IPInformationManager: - assert thl_web_rw.dsn.path - assert "/unittest-" in thl_web_rw.dsn.path - - from generalresearch.managers.thl.ipinfo import IPInformationManager - - return IPInformationManager(pg_config=thl_web_rw) - - -@pytest.fixture(scope="session") def ip_record_manager( thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig ) -> IPRecordManager: diff --git a/test_utils/managers/thl/conftest.py b/test_utils/managers/thl/conftest.py index 18a31e2..8ca4383 100644 --- a/test_utils/managers/thl/conftest.py +++ b/test_utils/managers/thl/conftest.py @@ -23,6 +23,10 @@ if TYPE_CHECKING: from generalresearch.config import GRLBaseSettings from generalresearch.managers.thl.buyer import BuyerManager from generalresearch.managers.thl.category import CategoryManager + from generalresearch.managers.thl.ipinfo import ( + IPGeonameManager, + IPInformationManager, + ) from generalresearch.managers.thl.payout import ( BrokerageProductPayoutEventManager, BusinessPayoutEventManager, @@ -40,6 +44,10 @@ if TYPE_CHECKING: from generalresearch.managers.thl.user_manager.user_metadata_manager import ( UserMetadataManager, ) + from generalresearch.managers.thl.userhealth import ( + AuditLogManager, + IPRecordManager, + ) from generalresearch.managers.thl.wall import ( WallCacheManager, WallManager, @@ -153,6 +161,13 @@ def brokerage_product_payout_event_manager( ) +@pytest.fixture() +def audit_log_manager(thl_web_rw: PostgresConfig) -> AuditLogManager: + from generalresearch.managers.thl.userhealth import AuditLogManager + + return AuditLogManager(pg_config=thl_web_rw) + + @pytest.fixture(scope="session") def business_payout_event_manager( thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig @@ -319,3 +334,41 @@ def surveypenalty_manager(thl_redis_config: RedisConfig): from generalresearch.managers.thl.survey_penalty import SurveyPenaltyManager return SurveyPenaltyManager(redis_config=thl_redis_config) + + +# --- IP Geolocation --- + + +@pytest.fixture +def ip_geoname_manager(thl_web_rw: PostgresConfig) -> IPGeonameManager: + from generalresearch.managers.thl.ipinfo import IPGeonameManager + + return IPGeonameManager(pg_config=thl_web_rw) + + +# --- IP Information --- + + +@pytest.fixture(scope="session") +def ip_information_manager(thl_web_rw: PostgresConfig) -> IPInformationManager: + assert thl_web_rw.dsn.path + assert "/unittest-" in thl_web_rw.dsn.path + + from generalresearch.managers.thl.ipinfo import IPInformationManager + + return IPInformationManager(pg_config=thl_web_rw) + + +# --- IP Record --- + + +@pytest.fixture(scope="session") +def ip_record_manager( + thl_web_rw: PostgresConfig, thl_redis_config: RedisConfig +) -> IPRecordManager: + assert thl_web_rw.dsn.path + assert "/unittest-" in thl_web_rw.dsn.path + + from generalresearch.managers.thl.userhealth import IPRecordManager + + return IPRecordManager(pg_config=thl_web_rw, redis_config=thl_redis_config) diff --git a/test_utils/models/conftest.py b/test_utils/models/conftest.py index d71593f..d5c9a71 100644 --- a/test_utils/models/conftest.py +++ b/test_utils/models/conftest.py @@ -10,6 +10,7 @@ from uuid import uuid4 import pytest from pydantic import AwareDatetime, PositiveInt +from pytest import FixtureRequest from pytest import FixtureRequest as Request from generalresearch.models.definitions import Source @@ -50,8 +51,6 @@ if TYPE_CHECKING: ) from generalresearch.models.thl.session import Session, Wall from generalresearch.models.thl.user import User - from generalresearch.models.thl.user_iphistory import IPRecord - from generalresearch.models.thl.userhealth import AuditLog, AuditLogLevel from generalresearch.pg_helper import PostgresConfig # === THL === @@ -59,15 +58,15 @@ if TYPE_CHECKING: @pytest.fixture def user( - request, - product_manager: ProductManager, + request: FixtureRequest, user_manager: UserManager, thl_web_rr: PostgresConfig, + product_factory: Callable[..., Product], ) -> User: product = getattr(request, "product", None) if product is None: - product = product_manager.create_dummy() + product = product_factory() u = user_manager.create_dummy(product_id=product.id) u.prefetch_product(pg_config=thl_web_rr) @@ -309,31 +308,35 @@ def payout_config(request: Request) -> PayoutConfig: @pytest.fixture def product_user_wallet_yes( - payout_config: PayoutConfig, product_manager: ProductManager + product_factory: Callable[..., Product], + payout_config: PayoutConfig, + product_manager: ProductManager, ) -> Product: from generalresearch.models.thl.product import UserWalletConfig - return product_manager.create_dummy( + return product_factory( payout_config=payout_config, user_wallet_config=UserWalletConfig(enabled=True) ) @pytest.fixture -def product_user_wallet_no(product_manager: ProductManager) -> Product: +def product_user_wallet_no( + product_factory: Callable[..., Product], product_manager: ProductManager +) -> Product: from generalresearch.models.thl.product import UserWalletConfig - return product_manager.create_dummy( - user_wallet_config=UserWalletConfig(enabled=False) - ) + return product_factory(user_wallet_config=UserWalletConfig(enabled=False)) @pytest.fixture def product_amt_true( - product_manager: ProductManager, payout_config: PayoutConfig + product_factory: Callable[..., Product], + product_manager: ProductManager, + payout_config: PayoutConfig, ) -> Product: from generalresearch.models.thl.product import UserWalletConfig - return product_manager.create_dummy( + return product_factory( user_wallet_config=UserWalletConfig(amt=True, enabled=True), payout_config=payout_config, ) @@ -370,84 +373,6 @@ def bp_payout_factory( return _inner -@pytest.fixture -def audit_log(audit_log_manager: AuditLogManager, user: User) -> AuditLog: - - return audit_log_manager.create_dummy(user_id=user.user_id) - - -@pytest.fixture -def audit_log_factory( - audit_log_manager: AuditLogManager, -) -> Callable[..., AuditLog]: - - def _inner( - user_id: PositiveInt, - level: AuditLogLevel | None = None, - event_type: str | None = None, - event_msg: str | None = None, - event_value: float | None = None, - ) -> AuditLog: - return audit_log_manager.create_dummy( - user_id=user_id, - level=level, - event_type=event_type, - event_msg=event_msg, - event_value=event_value, - ) - - return _inner - - -@pytest.fixture -def ip_geoname(ip_geoname_manager: IPGeonameManager) -> IPGeoname: - return ip_geoname_manager.create_dummy() - - -@pytest.fixture -def ip_information( - ip_information_manager: IPInformationManager, ip_geoname: IPGeoname -) -> IPInformation: - return ip_information_manager.create_dummy( - geoname_id=ip_geoname.geoname_id, country_iso=ip_geoname.country_iso - ) - - -@pytest.fixture -def ip_information_factory( - ip_information_manager: IPInformationManager, -) -> Callable[..., IPInformation]: - - def _inner(ip: str, geoname: IPGeoname, **kwargs) -> IPInformation: - return ip_information_manager.create_dummy( - ip=ip, - geoname_id=geoname.geoname_id, - country_iso=geoname.country_iso, - **kwargs, - ) - - return _inner - - -@pytest.fixture -def ip_record( - ip_record_manager: IPRecordManager, ip_geoname: IPGeoname, user: User -) -> IPRecord: - - return ip_record_manager.create_dummy(user_id=user.user_id) - - -@pytest.fixture -def ip_record_factory( - ip_record_manager: IPRecordManager, user: User -) -> Callable[..., IPRecord]: - - def _inner(user_id: PositiveInt, ip: str | None = None) -> IPRecord: - return ip_record_manager.create_dummy(user_id=user_id, ip=ip) - - return _inner - - @pytest.fixture(scope="session") def buyer(buyer_manager: BuyerManager) -> Buyer: buyer_code = uuid4().hex diff --git a/test_utils/models/thl/conftest.py b/test_utils/models/thl/conftest.py index 5826f0d..14f8f36 100644 --- a/test_utils/models/thl/conftest.py +++ b/test_utils/models/thl/conftest.py @@ -62,6 +62,7 @@ if TYPE_CHECKING: from generalresearch.models.thl.user_iphistory import IPRecord from generalresearch.models.thl.userhealth import AuditLog from generalresearch.models.thl.wallet.cashout_method import CashMailOrderData + from generalresearch.pg_helper import PostgresConfig fake = faker.Faker() @@ -72,30 +73,6 @@ def wall_status() -> Status: @pytest.fixture -def user_factory(user_manager: UserManager) -> Callable[..., User]: - - def _inner( - # --- Create dummy "optional" --- # - product_user_id: str | None = None, - # --- Optional --- # - product_id: UUIDStr | None = None, - product: Product | None = None, - created: datetime | None = None, - ) -> User: - - product_user_id = product_user_id or uuid4().hex - - return user_manager.create_user( - product_user_id=product_user_id, - product_id=product_id, - product=product, - created=created, - ) - - return _inner - - -@pytest.fixture def wall_factory( wall_manager: WallManager, session_factory: Session ) -> Callable[..., Wall]: @@ -155,9 +132,10 @@ def product_factory(product_manager: ProductManager) -> Callable[..., Product]: def _inner( save: bool = True, team: Team | None = None, + team_id: UUIDStr | None = None, business: Business | None = None, - product_id: UUIDStr | None = None, business_id: UUIDStr | None = None, + product_id: UUIDStr | None = None, name: str | None = None, redirect_url: str | None = None, harmonizer_domain: str | None = None, @@ -174,8 +152,10 @@ def product_factory(product_manager: ProductManager) -> Callable[..., Product]: product_id = product_id if product_id else uuid4().hex - team_id = team.uuid if team else uuid4().hex - business_id = business.uuid if business else uuid4().hex + team_id = (team.uuid if team else None) or team_id or uuid4().hex + business_id = ( + (business.uuid if business else None) or business_id or uuid4().hex + ) name = name if name else f"name-{product_id[:12]}" redirect_url = redirect_url if redirect_url else "https://www.example.com/" @@ -256,9 +236,12 @@ def session_factory(session_manager: SessionManager): @pytest.fixture -def ipgeoname_factory(ipgeoname_manager: IPGeonameManager) -> Callable[..., IPGeoname]: +def ip_geoname_factory( + ip_geoname_manager: IPGeonameManager, +) -> Callable[..., IPGeoname]: def _inner( + save: bool, geoname_id: PositiveInt | None = None, continent_code: str | None = None, continent_name: str | None = None, @@ -273,31 +256,47 @@ def ipgeoname_factory(ipgeoname_manager: IPGeonameManager) -> Callable[..., IPGe time_zone: str | None = None, is_in_european_union: bool | None = None, ) -> IPGeoname: - - return ipgeoname_manager.create( - geoname_id=geoname_id or randint(1, 999_999_999), - continent_code=continent_code or "na", - continent_name=continent_name or "North America", - country_iso=country_iso or "us", - country_name=country_name or "United States", - subdivision_1_iso=subdivision_1_iso or "fl", - subdivision_1_name=subdivision_1_name or "Florida", - subdivision_2_iso=subdivision_2_iso, - subdivision_2_name=subdivision_2_name, - city_name=city_name, - metro_code=metro_code, - time_zone=time_zone, - is_in_european_union=is_in_european_union, - ) + if save: + return ip_geoname_manager.create( + geoname_id=geoname_id or randint(1, 999_999_999), + continent_code=continent_code or "na", + continent_name=continent_name or "North America", + country_iso=country_iso or "us", + country_name=country_name or "United States", + subdivision_1_iso=subdivision_1_iso or "fl", + subdivision_1_name=subdivision_1_name or "Florida", + subdivision_2_iso=subdivision_2_iso, + subdivision_2_name=subdivision_2_name, + city_name=city_name, + metro_code=metro_code, + time_zone=time_zone, + is_in_european_union=is_in_european_union, + ) + else: + raise ValueError("Unsaved IPGeoname not yet supported") return _inner -def ipinformation_factory( +@pytest.fixture() +def ip_geoname(ip_geoname_factory: Callable[..., IPGeoname]) -> IPGeoname: + return ip_geoname_factory(save=True) + + +@pytest.fixture() +def unsaved_ip_geoname(ip_geoname_factory: Callable[..., IPGeoname]) -> IPGeoname: + return ip_geoname_factory(save=True) + + +# --- IP Information --- + + +def ip_information_factory( ipinformation_manager: IPInformationManager, ) -> Callable[..., IPInformation]: def _inner( + save: bool = True, ip: IPvAnyAddressStr | None = None, geoname_id: PositiveInt | None = None, country_iso: str | None = None, @@ -324,36 +323,184 @@ def ipinformation_factory( accuracy_radius: int | None = None, ) -> IPInformation: - return ipinformation_manager.create( - ip=ip or fake.ipv4_public(), - geoname_id=geoname_id, - country_iso=country_iso or fake.country_code(), - registered_country_iso=registered_country_iso, - is_anonymous=is_anonymous, - is_anonymous_vpn=is_anonymous_vpn, - is_hosting_provider=is_hosting_provider, - is_public_proxy=is_public_proxy, - is_tor_exit_node=is_tor_exit_node, - is_residential_proxy=is_residential_proxy, - autonomous_system_number=autonomous_system_number, - autonomous_system_organization=autonomous_system_organization, - domain=domain, - isp=isp, - mobile_country_code=mobile_country_code, - mobile_network_code=mobile_network_code, - network=network, - organization=organization, - static_ip_score=static_ip_score, - user_type=user_type, - postal_code=postal_code, - latitude=latitude, - longitude=longitude, - accuracy_radius=accuracy_radius, - ) + if save: + return ipinformation_manager.create( + ip=ip or fake.ipv4_public(), + geoname_id=geoname_id, + country_iso=country_iso or fake.country_code(), + registered_country_iso=registered_country_iso, + is_anonymous=is_anonymous, + is_anonymous_vpn=is_anonymous_vpn, + is_hosting_provider=is_hosting_provider, + is_public_proxy=is_public_proxy, + is_tor_exit_node=is_tor_exit_node, + is_residential_proxy=is_residential_proxy, + autonomous_system_number=autonomous_system_number, + autonomous_system_organization=autonomous_system_organization, + domain=domain, + isp=isp, + mobile_country_code=mobile_country_code, + mobile_network_code=mobile_network_code, + network=network, + organization=organization, + static_ip_score=static_ip_score, + user_type=user_type, + postal_code=postal_code, + latitude=latitude, + longitude=longitude, + accuracy_radius=accuracy_radius, + ) + else: + raise ValueError("Unsaved IP Information not supported yet") + + return _inner + + +@pytest.fixture +def ip_information( + ip_information_factory: Callable[..., IPInformation], +) -> IPInformation: + return ip_information_factory(save=True) + + +@pytest.fixture +def unsaved_ip_information( + ip_information_factory: Callable[..., IPInformation], +) -> IPInformation: + return ip_information_factory(save=False) + + +# --- IP Record --- + + +@pytest.fixture +def ip_record_factory( + ip_record_manager: IPRecordManager, user: User +) -> Callable[..., IPRecord]: + # return ip_record_manager.create_dummy(user_id=user.user_id) + + # def create_dummy( + # self, + # user_id: PositiveInt, + # ip: IPvAnyAddressStr | None = None, + # forwarded_ip1: IPvAnyAddressStr | None = None, + # forwarded_ip2: IPvAnyAddressStr | None = None, + # forwarded_ip3: IPvAnyAddressStr | None = None, + # forwarded_ip4: IPvAnyAddressStr | None = None, + # forwarded_ip5: IPvAnyAddressStr | None = None, + # forwarded_ip6: IPvAnyAddressStr | None = None, + # ) -> IPRecord: + # return self.create( + # user_id=user_id, + # ip=ip or fake.ipv4_public(), + # forwarded_ip1=(forwarded_ip1 or fake.ipv4_public()), + # forwarded_ip2=(forwarded_ip2 or fake.ipv6() if random() < 0.5 else None), + # forwarded_ip3=( + # forwarded_ip3 or fake.ipv4_public() if random() < 0.25 else None + # ), + # forwarded_ip4=forwarded_ip4, + # forwarded_ip5=forwarded_ip5, + # forwarded_ip6=forwarded_ip6, + # ) + + def _inner( + user_id: PositiveInt, save: bool = True, ip: str | None = None + ) -> IPRecord: + if save: + return ip_record_manager.create_dummy(user_id=user_id, ip=ip) + else: + raise ValueError("Unsaved IP Record not supported") return _inner +@pytest.fixture() +def ip_record( + ip_record_manager: IPRecordManager, ip_geoname: IPGeoname, user: User +) -> IPRecord: + return ip_record_factory(save=True) + + +@pytest.fixture() +def unsaved_ip_record(ip_record_factory: Callable[..., IPRecord]) -> IPRecord: + return ip_record_factory(save=False) + + +# --- User --- + + +@pytest.fixture() +def user_factory( + user_manager: UserManager, thl_web_rr: PostgresConfig +) -> Callable[..., User]: + + def _inner( + save: bool = True, + # --- Create dummy "optional" --- # + product_user_id: str | None = None, + # --- Optional --- # + product_id: UUIDStr | None = None, + product: Product | None = None, + created: datetime | None = None, + ) -> User: + if save: + if product is None: + product = product_factory() + + product_user_id = product_user_id or uuid4().hex + + u = user_manager.create_user( + product_user_id=product_user_id, + product_id=product_id, + product=product, + created=created, + ) + + u = user_manager.create_dummy(product=product, created=created) + + u.prefetch_product(pg_config=thl_web_rr) + return u + + else: + raise ValueError("Unsaved User not supported") + + return _inner + + +@pytest.fixture() +def user( + user_factory: Callable[..., User], +) -> User: + return user_factory(save=True) + + +@pytest.fixture() +def unsaved_user( + user_factory: Callable[..., User], +) -> User: + return user_factory(save=False) + + +@pytest.fixture +def user_with_wallet( + user_factory: Callable[..., User], + product_user_wallet_yes: Product, +) -> User: + # A user on a product with user wallet enabled, but they have no money + return user_factory(save=True, product=product_user_wallet_yes) + + +@pytest.fixture +def user_with_wallet_amt( + user_factory: Callable[..., User], product_amt_true: Product +) -> User: + # A user on a product with user wallet enabled, on AMT, but they have no money + return user_factory(save=True, product=product_amt_true) + + +# --- User Payout --- + + @pytest.fixture def user_payout_event_factory( user_payout_event_manager: UserPayoutEventManager, @@ -437,11 +584,11 @@ def iprecord_factory(iprecord_manager: IPRecordManager) -> Callable[..., IPRecor return _inner -# class AuditLogManager(PostgresManager): +# --- Audit Log Manager --- -@pytest.fixture -def auditlog_factory(audit_log_manager: AuditLogManager): +@pytest.fixture() +def audit_log_factory(audit_log_manager: AuditLogManager) -> Callable[..., AuditLog]: def _inner( user_id: PositiveInt, @@ -468,6 +615,19 @@ def auditlog_factory(audit_log_manager: AuditLogManager): return _inner +@pytest.fixture() +def audit_log(auditlog_factory: Callable[..., AuditLog]) -> AuditLog: + return auditlog_factory(save=True) + + +@pytest.fixture() +def unsaved_audit_log(auditlog_factory: Callable[..., AuditLog]) -> AuditLog: + return auditlog_factory(save=False) + + +# --- --- + + @pytest.fixture(scope="session") def profiling_info_json() -> str: return ( diff --git a/tests/grliq/managers/test_forensic_data.py b/tests/grliq/managers/test_forensic_data.py index e4854e8..1b83757 100644 --- a/tests/grliq/managers/test_forensic_data.py +++ b/tests/grliq/managers/test_forensic_data.py @@ -1,5 +1,6 @@ from __future__ import annotations +from collections.abc import Callable from datetime import timedelta from typing import TYPE_CHECKING from uuid import uuid4 @@ -16,6 +17,8 @@ from generalresearch.grliq.models.forensic_result import ( if TYPE_CHECKING: from generalresearch.grliq.managers.forensic_data import ( GrlIqDataManager, + ) + from generalresearch.grliq.managers.forensic_events import ( GrlIqEventManager, ) from generalresearch.models.thl.product import Product @@ -28,10 +31,13 @@ except ImportError: class TestGrlIqDataManager: - def test_create_dummy(self, grliq_dm: GrlIqDataManager): + def test_create_dummy( + self, + grliq_data_factory: Callable[..., GrlIqData], + ): from generalresearch.grliq.models.forensic_data import GrlIqData - gd1: GrlIqData = grliq_dm.create_dummy(is_attempt_allowed=True) + gd1: GrlIqData = grliq_data_factory(is_attempt_allowed=True) assert isinstance(gd1, GrlIqData) assert isinstance(gd1.results, GrlIqCheckerResults) @@ -119,7 +125,9 @@ class TestGrlIqDataManager: class TestForensicDataGetAndFilter: - def test_events(self, grliq_dm: GrlIqDataManager): + def test_events( + self, grliq_dm: GrlIqDataManager, grliq_data_factory: Callable[..., GrlIqData] + ): """If load_events=True, the events and mouse_events attributes should be an array no matter what. An empty array means that the events were loaded, but there were no events available. @@ -129,7 +137,7 @@ class TestForensicDataGetAndFilter: """ # Load Events == False forensic_uuid = uuid4().hex - grliq_dm.create_dummy(is_attempt_allowed=True, uuid=forensic_uuid) + grliq_data_factory(is_attempt_allowed=True, uuid=forensic_uuid) instance = grliq_dm.filter_data(uuids=[forensic_uuid])[0] assert isinstance(instance, GrlIqData) @@ -144,41 +152,53 @@ class TestForensicDataGetAndFilter: assert len(instance.events) == 0 assert len(instance.mouse_events) == 0 - def test_timing(self, grliq_dm: GrlIqDataManager, grliq_em: GrlIqEventManager): + def test_timing( + self, + grliq_data_factory: Callable[..., GrlIqData], + grliq_data_manager: GrlIqDataManager, + grliq_event_manager: GrlIqEventManager, + ): forensic_uuid = uuid4().hex - grliq_dm.create_dummy(is_attempt_allowed=True, uuid=forensic_uuid) + grliq_data_factory(is_attempt_allowed=True, uuid=forensic_uuid) - instance = grliq_dm.filter_data(uuids=[forensic_uuid])[0] + instance = grliq_data_manager.filter_data(uuids=[forensic_uuid])[0] - grliq_em.update_or_create_timing( + grliq_event_manager.update_or_create_timing( session_uuid=instance.mid, timing_data=TimingData( client_rtts=[100, 200, 150], server_rtts=[150, 120, 120] ), ) - instance = grliq_dm.get_data(forensic_uuid=forensic_uuid, load_events=True) + instance = grliq_data_manager.get_data( + forensic_uuid=forensic_uuid, load_events=True + ) assert isinstance(instance, GrlIqData) assert isinstance(instance.events, list) assert isinstance(instance.mouse_events, list) assert isinstance(instance.timing_data, TimingData) def test_events_events( - self, grliq_dm: GrlIqDataManager, grliq_em: GrlIqEventManager + self, + grliq_data_factory: Callable[..., GrlIqData], + grliq_data_manager: GrlIqDataManager, + grliq_event_manager: GrlIqEventManager, ): forensic_uuid = uuid4().hex - grliq_dm.create_dummy(is_attempt_allowed=True, uuid=forensic_uuid) + grliq_data_factory(is_attempt_allowed=True, uuid=forensic_uuid) - instance = grliq_dm.filter_data(uuids=[forensic_uuid])[0] + instance = grliq_data_manager.filter_data(uuids=[forensic_uuid])[0] - grliq_em.update_or_create_events( + grliq_event_manager.update_or_create_events( session_uuid=instance.mid, events=[{"a": "b"}], mouse_events=[], event_start=instance.created_at, event_end=instance.created_at + timedelta(minutes=1), ) - instance = grliq_dm.get_data(forensic_uuid=forensic_uuid, load_events=True) + instance = grliq_data_manager.get_data( + forensic_uuid=forensic_uuid, load_events=True + ) assert isinstance(instance, GrlIqData) assert isinstance(instance.events, list) assert isinstance(instance.mouse_events, list) @@ -189,11 +209,16 @@ class TestForensicDataGetAndFilter: assert len(instance.keyboard_events) == 0 def test_events_click( - self, grliq_dm: GrlIqDataManager, grliq_em: GrlIqEventManager + self, + grliq_data_factory: Callable[..., GrlIqData], + grliq_data_manager: GrlIqDataManager, + grliq_event_manager: GrlIqEventManager, ): forensic_uuid = uuid4().hex - grliq_dm.create_dummy(is_attempt_allowed=True, uuid=forensic_uuid) - instance = grliq_dm.get_data(forensic_uuid=forensic_uuid, load_events=True) + grliq_data_factory(is_attempt_allowed=True, uuid=forensic_uuid) + instance = grliq_data_manager.get_data( + forensic_uuid=forensic_uuid, load_events=True + ) click_event = { "type": "click", @@ -203,14 +228,16 @@ class TestForensicDataGetAndFilter: "pointerType": "mouse", } me = MouseEvent.from_dict(click_event) - grliq_em.update_or_create_events( + grliq_event_manager.update_or_create_events( session_uuid=instance.mid, events=[click_event], mouse_events=[], event_start=instance.created_at, event_end=instance.created_at + timedelta(minutes=1), ) - instance = grliq_dm.get_data(forensic_uuid=forensic_uuid, load_events=True) + instance = grliq_data_manager.get_data( + forensic_uuid=forensic_uuid, load_events=True + ) assert isinstance(instance, GrlIqData) assert isinstance(instance.events, list) assert isinstance(instance.mouse_events, list) diff --git a/tests/grliq/managers/test_forensic_results.py b/tests/grliq/managers/test_forensic_results.py index a030451..86834d0 100644 --- a/tests/grliq/managers/test_forensic_results.py +++ b/tests/grliq/managers/test_forensic_results.py @@ -1,18 +1,21 @@ from __future__ import annotations +from collections.abc import Callable from typing import TYPE_CHECKING if TYPE_CHECKING: - from generalresearch.grliq.managers.forensic_data import GrlIqDataManager from generalresearch.grliq.managers.forensic_results import ( GrlIqCategoryResultsReader, ) + from generalresearch.grliq.models.forensic_data import GrlIqData class TestGrlIqCategoryResultsReader: def test_filter_category_results( - self, grliq_dm: GrlIqDataManager, grliq_crr: GrlIqCategoryResultsReader + self, + grliq_data_factory: Callable[..., GrlIqData], + grliq_crr: GrlIqCategoryResultsReader, ): from generalresearch.grliq.models.forensic_result import ( GrlIqForensicCategoryResult, @@ -20,8 +23,8 @@ class TestGrlIqCategoryResultsReader: ) # this is just testing that it doesn't fail - grliq_dm.create_dummy(is_attempt_allowed=True) - grliq_dm.create_dummy(is_attempt_allowed=True) + grliq_data_factory(is_attempt_allowed=True) + grliq_data_factory(is_attempt_allowed=True) res = grliq_crr.filter_category_results(limit=2, phase=Phase.OFFERWALL_ENTER)[0] assert res.get("category_result") diff --git a/tests/managers/gr/test_business.py b/tests/managers/gr/test_business.py index 0d5b0d5..6a930b4 100644 --- a/tests/managers/gr/test_business.py +++ b/tests/managers/gr/test_business.py @@ -76,30 +76,30 @@ class TestBusinessManager: assert isinstance(instance, Business) assert isinstance(instance.id, int) - def test_get_or_create(self, business_manager: BusinessManager): + def test_get_or_create(self, gr_business_manager: BusinessManager): uuid_key = uuid4().hex - assert business_manager.get_by_uuid(business_uuid=uuid_key) is None + assert gr_business_manager.get_by_uuid(business_uuid=uuid_key) is None - instance = business_manager.get_or_create( + instance = gr_business_manager.get_or_create( uuid=uuid_key, name=f"name-{uuid4().hex[:6]}", ) - res = business_manager.get_by_uuid(business_uuid=uuid_key) + res = gr_business_manager.get_by_uuid(business_uuid=uuid_key) assert isinstance(res, Business) assert res.id == instance.id def test_get_all( self, - business_manager: BusinessManager, + gr_business_manager: BusinessManager, gr_business_factory: Callable[..., Business], ): - res1 = business_manager.get_all() + res1 = gr_business_manager.get_all() assert isinstance(res1, list) gr_business_factory() - res2 = business_manager.get_all() + res2 = gr_business_manager.get_all() assert len(res1) == len(res2) - 1 @pytest.mark.skip(reason="TODO") @@ -108,42 +108,42 @@ class TestBusinessManager: def test_get_by_user_id( self, - business_manager: BusinessManager, + gr_business_manager: BusinessManager, gr_user: GRUser, team_manager: TeamManager, membership_manager: MembershipManager, gr_business_factory: Callable[..., Business], gr_team_factory: Callable[..., Team], ): - res = business_manager.get_by_user_id(user_id=gr_user.id) + res = gr_business_manager.get_by_user_id(user_id=gr_user.id) assert len(res) == 0 # Create a business: Business, but don't add it to anything b1 = gr_business_factory() - res = business_manager.get_by_user_id(user_id=gr_user.id) + res = gr_business_manager.get_by_user_id(user_id=gr_user.id) assert len(res) == 0 # Create a Team, but don't create any Memberships t1 = gr_team_factory() - res = business_manager.get_by_user_id(user_id=gr_user.id) + res = gr_business_manager.get_by_user_id(user_id=gr_user.id) assert len(res) == 0 # Create a Membership for the gr_user to the Team... but it doesn't # matter because the Team doesn't have any Business yet _ = membership_manager.create(team=t1, gr_user=gr_user) - res = business_manager.get_by_user_id(user_id=gr_user.id) + res = gr_business_manager.get_by_user_id(user_id=gr_user.id) assert len(res) == 0 # Add the Business to the Team... now the Business should be available # to the gr_user team_manager.add_business(team=t1, business=b1) - res = business_manager.get_by_user_id(user_id=gr_user.id) + res = gr_business_manager.get_by_user_id(user_id=gr_user.id) assert len(res) == 1 # Add another Business to the Team! b2 = gr_business_factory() team_manager.add_business(team=t1, business=b2) - res = business_manager.get_by_user_id(user_id=gr_user.id) + res = gr_business_manager.get_by_user_id(user_id=gr_user.id) assert len(res) == 2 @pytest.mark.skip(reason="TODO") @@ -151,14 +151,16 @@ class TestBusinessManager: pass def test_get_by_uuid( - self, gr_business: Business, business_manager: BusinessManager + self, gr_business: Business, gr_business_manager: BusinessManager ): - instance = business_manager.get_by_uuid(business_uuid=gr_business.uuid) + instance = gr_business_manager.get_by_uuid(business_uuid=gr_business.uuid) assert isinstance(instance, Business) assert gr_business.id == instance.id - def test_get_by_id(self, gr_business: Business, business_manager: BusinessManager): - instance = business_manager.get_by_id(business_id=gr_business.id) + def test_get_by_id( + self, gr_business: Business, gr_business_manager: BusinessManager + ): + instance = gr_business_manager.get_by_id(business_id=gr_business.id) assert isinstance(instance, Business) assert gr_business.uuid == instance.uuid diff --git a/tests/managers/thl/test_contest/test_milestone.py b/tests/managers/thl/test_contest/test_milestone.py index dbb2016..dab02e7 100644 --- a/tests/managers/thl/test_contest/test_milestone.py +++ b/tests/managers/thl/test_contest/test_milestone.py @@ -6,10 +6,10 @@ from typing import TYPE_CHECKING from generalresearch.models.thl.contest.definitions import ( ContestEndReason, + ContestEntryTrigger, ContestStatus, ) from generalresearch.models.thl.contest.milestone import ( - ContestEntryTrigger, MilestoneContest, MilestoneUserView, ) diff --git a/tests/managers/thl/test_contest/test_raffle.py b/tests/managers/thl/test_contest/test_raffle.py index 7803952..0b2b852 100644 --- a/tests/managers/thl/test_contest/test_raffle.py +++ b/tests/managers/thl/test_contest/test_raffle.py @@ -17,17 +17,17 @@ from generalresearch.models.thl.contest import ( ContestEntryRule, ContestPrize, ) +from generalresearch.models.thl.contest.contest_entry import ( + ContestEntry, + ContestEntryType, +) from generalresearch.models.thl.contest.definitions import ( ContestEndReason, ContestPrizeKind, ContestStatus, ) from generalresearch.models.thl.contest.exceptions import ContestError -from generalresearch.models.thl.contest.raffle import ( - ContestEntry, - ContestEntryType, - RaffleContest, -) +from generalresearch.models.thl.contest.raffle import RaffleContest if TYPE_CHECKING: from generalresearch.managers.thl.contest_manager import ContestManager diff --git a/tests/managers/thl/test_ipinfo.py b/tests/managers/thl/test_ipinfo.py index 6954163..47b1712 100644 --- a/tests/managers/thl/test_ipinfo.py +++ b/tests/managers/thl/test_ipinfo.py @@ -31,14 +31,16 @@ class TestIPGeonameManager: assert isinstance(instance, IPGeonameManager) assert isinstance(ip_geoname_manager, IPGeonameManager) - def test_create(self, ip_geoname_manager: IPGeonameManager): - - instance = ip_geoname_manager.create_dummy() + def test_create( + self, + ip_geoname_factory: Callable[..., IPGeoname], + ip_geoname_manager: IPGeonameManager, + ): + instance = ip_geoname_factory() assert isinstance(instance, IPGeoname) res = ip_geoname_manager.fetch_geoname_ids(filter_ids=[instance.geoname_id]) - assert res[0].model_dump_json() == instance.model_dump_json() @@ -51,13 +53,15 @@ class TestIPInformationManager: assert isinstance(instance, IPInformationManager) assert isinstance(ip_information_manager, IPInformationManager) - def test_create(self, ip_information_manager: IPInformationManager): - instance = ip_information_manager.create_dummy() - + def test_create( + self, + ip_geoname_factory: Callable[..., IPGeoname], + ip_information_manager: IPInformationManager, + ): + instance = ip_geoname_factory() assert isinstance(instance, IPInformation) res = ip_information_manager.fetch_ip_information(filter_ips=[instance.ip]) - assert res[0].model_dump_json() == instance.model_dump_json() def test_prefetch_geoname( diff --git a/tests/managers/thl/test_ledger/test_lm_accounts.py b/tests/managers/thl/test_ledger/test_lm_accounts.py index 57a2261..3af10e7 100644 --- a/tests/managers/thl/test_ledger/test_lm_accounts.py +++ b/tests/managers/thl/test_ledger/test_lm_accounts.py @@ -14,8 +14,10 @@ from generalresearch.managers.thl.ledger_manager.exceptions import ( LedgerAccountDoesntExistError, ) from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager -from generalresearch.models.custom_types import AccountType, Direction, UUIDStr +from generalresearch.models.custom_types import UUIDStr from generalresearch.models.thl.ledger import ( + AccountType, + Direction, LedgerAccount, LedgerEntry, ) diff --git a/tests/managers/thl/test_ledger/test_lm_tx_locks.py b/tests/managers/thl/test_ledger/test_lm_tx_locks.py index 9ecc1bc..166598e 100644 --- a/tests/managers/thl/test_ledger/test_lm_tx_locks.py +++ b/tests/managers/thl/test_ledger/test_lm_tx_locks.py @@ -4,10 +4,10 @@ import logging from collections.abc import Callable, Generator from datetime import UTC, datetime, timedelta from decimal import Decimal -from logging import LogCaptureFixture from typing import TYPE_CHECKING import pytest +from pytest import LogCaptureFixture from generalresearch.managers.thl.ledger_manager.conditions import ( generate_condition_mp_payment, diff --git a/tests/managers/thl/test_ledger/test_thl_lm_tx.py b/tests/managers/thl/test_ledger/test_thl_lm_tx.py index b0484ae..cda88da 100644 --- a/tests/managers/thl/test_ledger/test_thl_lm_tx.py +++ b/tests/managers/thl/test_ledger/test_thl_lm_tx.py @@ -128,10 +128,11 @@ class TestThlLedgerTxManager: thl_ledger_manager: ThlLedgerManager, ledger_manager: LedgerManager, session_manager: SessionManager, + product_factory: Callable[..., Product], ): delete_ledger_db() create_main_accounts() - product = product_manager.create_dummy( + product = product_factory( payout_config=PayoutConfig( payout_transformation=PayoutTransformation( f="payout_transformation_amt" diff --git a/tests/managers/thl/test_ledger/test_wallet.py b/tests/managers/thl/test_ledger/test_wallet.py index 1ee9bf9..dc1feec 100644 --- a/tests/managers/thl/test_ledger/test_wallet.py +++ b/tests/managers/thl/test_ledger/test_wallet.py @@ -22,8 +22,10 @@ if TYPE_CHECKING: @pytest.fixture() -def schrute_product(product_manager: ProductManager) -> Product: - return product_manager.create_dummy( +def schrute_product( + product_factory: Callable[..., Product], product_manager: ProductManager +) -> Product: + return product_factory( user_wallet_config=UserWalletConfig(enabled=True, amt=False), payout_config=PayoutConfig( payout_transformation=PayoutTransformation( 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( diff --git a/tests/managers/thl/test_task_status.py b/tests/managers/thl/test_task_status.py index 9846ce0..4a401fa 100644 --- a/tests/managers/thl/test_task_status.py +++ b/tests/managers/thl/test_task_status.py @@ -40,18 +40,22 @@ finish3 = start3 + timedelta(minutes=5) @pytest.fixture(scope="session") -def bp1(product_manager: ProductManager) -> Product: +def bp1( + product_factory: Callable[..., Product], product_manager: ProductManager +) -> Product: # user wallet disabled, payout xform NULL - return product_manager.create_dummy( + return product_factory( user_wallet_config=UserWalletConfig(enabled=False), payout_config=PayoutConfig(), ) @pytest.fixture(scope="session") -def bp2(product_manager: ProductManager) -> Product: +def bp2( + product_factory: Callable[..., Product], product_manager: ProductManager +) -> Product: # user wallet disabled, payout xform 40% - return product_manager.create_dummy( + return product_factory( user_wallet_config=UserWalletConfig(enabled=False), payout_config=PayoutConfig( payout_transformation=PayoutTransformation( @@ -63,9 +67,11 @@ def bp2(product_manager: ProductManager) -> Product: @pytest.fixture(scope="session") -def bp3(product_manager: ProductManager) -> Product: +def bp3( + product_factory: Callable[..., Product], product_manager: ProductManager +) -> Product: # user wallet enabled, payout xform 50% - return product_manager.create_dummy( + return product_factory( user_wallet_config=UserWalletConfig(enabled=True), payout_config=PayoutConfig( payout_transformation=PayoutTransformation( diff --git a/tests/managers/thl/test_user_manager/test_base.py b/tests/managers/thl/test_user_manager/test_base.py index 4a9750e..c69f297 100644 --- a/tests/managers/thl/test_user_manager/test_base.py +++ b/tests/managers/thl/test_user_manager/test_base.py @@ -1,4 +1,5 @@ import logging +from collections.abc import Callable from datetime import UTC, datetime from random import randint from typing import TYPE_CHECKING @@ -7,9 +8,11 @@ from uuid import uuid4 import pytest from generalresearch.managers.thl.user_manager import ( - UserCreateNotAllowedError, get_bp_user_create_limit_hourly, ) +from generalresearch.managers.thl.user_manager.exceptions import ( + UserCreateNotAllowedError, +) from generalresearch.managers.thl.user_manager.mysql_user_manager import ( MysqlUserManager, ) @@ -152,11 +155,11 @@ class TestCreateUserManager: def test_create_user( self, - product_manager: ProductManager, + product_factory: Callable[..., Product], thl_web_rw: PostgresConfig, user_manager: UserManager, ): - product: Product = product_manager.create_dummy( + product: Product = product_factory( user_create_config=UserCreateConfig( min_hourly_create_limit=10, max_hourly_create_limit=69 ), @@ -195,11 +198,11 @@ class TestCreateUserManager: def test_create_user_integrity_error( self, - product_manager: ProductManager, user_manager: UserManager, + product_factory: Callable[..., Product], caplog, ): - product: Product = product_manager.create_dummy( + product: Product = product_factory( product_id=uuid4().hex, team_id=uuid4().hex, name=f"Test Product ID #{uuid4().hex[:6]}", @@ -241,10 +244,13 @@ class TestCreateUserManager: assert user1 == user2 def test_raise_allow_user_create( - self, product_manager: ProductManager, user_manager: UserManager + self, + product_manager: ProductManager, + user_manager: UserManager, + product_factory: Callable[..., Product], ): rand_num = randint(25, 200) - product: Product = product_manager.create_dummy( + product: Product = product_factory( product_id=uuid4().hex, team_id=uuid4().hex, name=f"Test Product ID #{uuid4().hex[:6]}", diff --git a/tests/managers/thl/test_wall_manager.py b/tests/managers/thl/test_wall_manager.py index 3215de8..58de7a2 100644 --- a/tests/managers/thl/test_wall_manager.py +++ b/tests/managers/thl/test_wall_manager.py @@ -10,7 +10,7 @@ import pytest from pydantic import PositiveInt from generalresearch.models.definitions import Source -from generalresearch.models.thl.session import ( +from generalresearch.models.thl.definitions import ( ReportValue, Status, StatusCode1, diff --git a/tests/models/gr/test_business.py b/tests/models/gr/test_business.py index 4e0b4e1..e942be5 100644 --- a/tests/models/gr/test_business.py +++ b/tests/models/gr/test_business.py @@ -551,7 +551,6 @@ class TestBusinessBalance: ledger_manager: LedgerManager, product_manager: ProductManager, start: datetime, - thl_web_rr: PostgresConfig, session_with_tx_factory: Callable[..., Session], delete_ledger_db: Callable[..., None], create_main_accounts: Callable[..., None], @@ -964,7 +963,6 @@ class TestBusinessBalance: ledger_manager: LedgerManager, product_manager: ProductManager, start: datetime, - thl_web_rr: PostgresConfig, payout_event_manager, session_with_tx_factory: Callable[..., None], delete_ledger_db: Callable[..., None], @@ -1194,7 +1192,6 @@ class TestBusinessMethods: def test_set_cache( self, gr_business: Business, - gr_db: PostgresConfig, thl_web_rr: PostgresConfig, client_no_amm: DaskClient, mnt_filepath: GRLDatasets, diff --git a/tests/models/gr/test_team.py b/tests/models/gr/test_team.py index a94d53f..0ca9b11 100644 --- a/tests/models/gr/test_team.py +++ b/tests/models/gr/test_team.py @@ -113,13 +113,13 @@ class TestTeam: assert gr_team.businesses is None - gr_team.prefetch_businesses(business_manager=gr_business_manager) + 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=gr_team, business=business) assert len(gr_team.businesses) == 0 - gr_team.prefetch_businesses(business_manager=gr_business_manager) + 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 == business.uuid @@ -163,12 +163,19 @@ class TestTeamMethods: 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, ): client = gr_redis_config.create_redis_client() assert client.get(name=gr_team.cache_key) is None gr_team.set_cache( - pg_config=gr_db, + product_manager=product_manager, + gr_user_manager=gr_user_manager, + gr_business_manager=gr_business_manager, + gr_membership_manager=gr_membership_manager, thl_web_rr=thl_web_rr, redis_config=gr_redis_config, client=client_no_amm, @@ -193,6 +200,10 @@ class TestTeamMethods: 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 @@ -200,7 +211,10 @@ class TestTeamMethods: membership_factory(team=gr_team, gr_user=gr_user) gr_team.set_cache( - pg_config=gr_db, + product_manager=product_manager, + gr_user_manager=gr_user_manager, + gr_business_manager=gr_business_manager, + gr_membership_manager=gr_membership_manager, thl_web_rr=thl_web_rr, redis_config=gr_redis_config, client=client_no_amm, @@ -239,6 +253,7 @@ class TestTeamMethods: mnt_filepath: GRLDatasets, mnt_gr_api_dir: Path, gr_team: Team, + product_manager: ProductManager, ): delete_df_collection(coll=wall_collection) @@ -267,7 +282,7 @@ class TestTeamMethods: ) gr_team.prebuild_enriched_session_parquet( - thl_pg_config=thl_web_rr, + product_manager=product_manager, ds=mnt_filepath, client=client_no_amm, mnt_gr_api=mnt_gr_api_dir, @@ -295,6 +310,7 @@ class TestTeamMethods: mnt_filepath: GRLDatasets, mnt_gr_api_dir: Path, gr_team: Team, + product_manager: ProductManager, ): delete_df_collection(coll=wall_collection) @@ -323,7 +339,7 @@ class TestTeamMethods: ) gr_team.prebuild_enriched_wall_parquet( - thl_pg_config=thl_web_rr, + product_manager=product_manager, ds=mnt_filepath, client=client_no_amm, mnt_gr_api=mnt_gr_api_dir, diff --git a/tests/models/thl/test_product.py b/tests/models/thl/test_product.py index 446b59f..f1050bb 100644 --- a/tests/models/thl/test_product.py +++ b/tests/models/thl/test_product.py @@ -64,9 +64,11 @@ 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) res = instance.model_dump_json() diff --git a/tests/models/thl/test_user.py b/tests/models/thl/test_user.py index bc941d4..68b413c 100644 --- a/tests/models/thl/test_user.py +++ b/tests/models/thl/test_user.py @@ -18,6 +18,7 @@ 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: @@ -621,12 +622,17 @@ class TestUserSerialization: class TestUserMethods: - def test_audit_log(self, user: User, audit_log_manager: AuditLogManager): + 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 |
