diff options
| author | Max Nanis | 2026-08-21 11:31:13 -0700 |
|---|---|---|
| committer | Max Nanis | 2026-08-21 11:31:13 -0700 |
| commit | ff538deb290c851364df85aa6ccc64bc3719ef64 (patch) | |
| tree | c04c3573ab6adab70971fef00d3420a5da5f9619 /test_utils/models | |
| parent | 4fe0f6b5e0f0c744902e4c3ab8940e23a6f8a2e1 (diff) | |
| parent | c4a44873540ca4c0a3ab19b9beef4cfc6e0252a7 (diff) | |
| download | generalresearch-ff538deb290c851364df85aa6ccc64bc3719ef64.tar.gz generalresearch-ff538deb290c851364df85aa6ccc64bc3719ef64.zip | |
Merge branch 'master' into dev
Diffstat (limited to 'test_utils/models')
21 files changed, 2075 insertions, 159 deletions
diff --git a/test_utils/models/conftest.py b/test_utils/models/conftest.py index 6a8e4cf..89f6f32 100644 --- a/test_utils/models/conftest.py +++ b/test_utils/models/conftest.py @@ -1,11 +1,14 @@ +from __future__ import annotations + from datetime import datetime, timedelta, timezone from decimal import Decimal from random import choice as randchoice from random import randint -from typing import TYPE_CHECKING, Callable, Dict, List, Optional +from typing import TYPE_CHECKING, Callable from uuid import uuid4 import pytest +from fastapi import Request from pydantic import AwareDatetime, PositiveInt from generalresearch.models import Source @@ -13,6 +16,7 @@ from generalresearch.models.thl.definitions import ( WALL_ALLOWED_STATUS_STATUS_CODE, Status, ) +from generalresearch.models.thl.survey.model import Buyer, Survey from generalresearch.pg_helper import PostgresConfig from generalresearch.redis_helper import RedisConfig @@ -56,7 +60,6 @@ if TYPE_CHECKING: Product, ) from generalresearch.models.thl.session import Session, Wall - from generalresearch.models.thl.survey.model import Buyer, Survey from generalresearch.models.thl.user import User from generalresearch.models.thl.user_iphistory import IPRecord from generalresearch.models.thl.userhealth import AuditLog, AuditLogLevel @@ -67,10 +70,10 @@ if TYPE_CHECKING: @pytest.fixture def user( request, - product_manager: "ProductManager", - user_manager: "UserManager", + product_manager: ProductManager, + user_manager: UserManager, thl_web_rr: PostgresConfig, -) -> "User": +) -> User: product = getattr(request, "product", None) if product is None: @@ -84,26 +87,27 @@ def user( @pytest.fixture def user_with_wallet( - request, user_factory: Callable[..., "User"], product_user_wallet_yes: "Product" -) -> "User": + 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(product=product_user_wallet_yes) @pytest.fixture def user_with_wallet_amt( - request, user_factory: Callable[..., "User"], product_amt_true: "Product" -) -> "User": + 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(product=product_amt_true) @pytest.fixture(scope="function") def user_factory( - user_manager: "UserManager", thl_web_rr: PostgresConfig -) -> Callable[..., "User"]: + user_manager: UserManager, thl_web_rr: PostgresConfig +) -> Callable[..., User]: - def _inner(product: "Product", created: Optional[datetime] = None) -> "User": + def _inner(product: Product, created: datetime | None = None) -> User: u = user_manager.create_dummy(product=product, created=created) u.prefetch_product(pg_config=thl_web_rr) @@ -113,11 +117,11 @@ def user_factory( @pytest.fixture -def wall_factory(wall_manager: "WallManager") -> Callable[..., "Wall"]: +def wall_factory(wall_manager: WallManager) -> Callable[..., Wall]: def _inner( - session: "Session", wall_status: "Status", req_cpi: Optional[Decimal] = None - ) -> "Wall": + session: Session, wall_status: Status, req_cpi: Decimal | None = None + ) -> Wall: assert session.started <= datetime.now( tz=timezone.utc @@ -153,9 +157,7 @@ def wall_factory(wall_manager: "WallManager") -> Callable[..., "Wall"]: @pytest.fixture -def wall( - session: "Session", user: "User", wall_manager: "WallManager" -) -> Optional["Wall"]: +def wall(session: Session, user: User, wall_manager: WallManager) -> Wall | None: from generalresearch.models.thl.task_status import StatusCode1 wall = wall_manager.create_dummy(session_id=session.id, user_id=user.user_id) @@ -170,25 +172,24 @@ def wall( @pytest.fixture def session_factory( - wall_factory: Callable[..., "Wall"], - session_manager: "SessionManager", - wall_manager: "WallManager", + session_manager: SessionManager, + wall_manager: WallManager, utc_hour_ago: datetime, -) -> Callable[..., "Session"]: +) -> Callable[..., Session]: from generalresearch.models.thl.session import Source def _inner( - user: "User", + user: User, # Wall details wall_count: int = 5, wall_req_cpi: Decimal = Decimal(".50"), - wall_req_cpis: Optional[List[Decimal]] = None, - wall_statuses: Optional[List[Status]] = None, + wall_req_cpis: list[Decimal] | None = None, + wall_statuses: list[Status] | None = None, wall_source: Source = Source.TESTING, # Session details final_status: Status = Status.COMPLETE, started: datetime = utc_hour_ago, - ) -> "Session": + ) -> Session: if wall_req_cpis: assert len(wall_req_cpis) == wall_count if wall_statuses: @@ -236,24 +237,24 @@ def session_factory( @pytest.fixture(scope="function") def finished_session_factory( - session_factory: Callable[..., "Session"], - session_manager: "SessionManager", + session_factory: Callable[..., Session], + session_manager: SessionManager, utc_hour_ago: datetime, -) -> Callable[..., "Session"]: +) -> Callable[..., Session]: from generalresearch.models.thl.session import Source def _inner( - user: "User", + user: User, # Wall details wall_count: int = 5, wall_req_cpi: Decimal = Decimal(".50"), - wall_req_cpis: Optional[List[Decimal]] = None, - wall_statuses: Optional[List[Status]] = None, + wall_req_cpis: list[Decimal] | None = None, + wall_statuses: list[Status] | None = None, wall_source: Source = Source.TESTING, # Session details final_status: Status = Status.COMPLETE, started: datetime = utc_hour_ago, - ) -> "Session": + ) -> Session: s: Session = session_factory( user=user, wall_count=wall_count, @@ -281,9 +282,8 @@ def finished_session_factory( @pytest.fixture def session( - user: "User", session_manager: "SessionManager", wall_manager: "WallManager" -) -> "Session": - from generalresearch.models.thl.session import Session, Wall + user: User, session_manager: SessionManager, wall_manager: WallManager +) -> Session: session: Session = session_manager.create_dummy(user=user, country_iso="us") wall: Wall = wall_manager.create_dummy( @@ -297,7 +297,7 @@ def session( @pytest.fixture -def product(request, product_manager: "ProductManager") -> "Product": +def product(request: Request, product_manager: ProductManager) -> Product: team = getattr(request, "team", None) business = getattr(request, "business", None) @@ -309,13 +309,13 @@ def product(request, product_manager: "ProductManager") -> "Product": @pytest.fixture -def product_factory(product_manager: "ProductManager") -> Callable[..., "Product"]: +def product_factory(product_manager: ProductManager) -> Callable[..., Product]: def _inner( - team: Optional["Team"] = None, - business: Optional["Business"] = None, + team: Team | None = None, + business: Business | None = None, commission_pct: Decimal = Decimal("0.05"), - ) -> "Product": + ) -> Product: return product_manager.create_dummy( team_id=team.uuid if team else None, business_id=business.uuid if business else None, @@ -326,7 +326,7 @@ def product_factory(product_manager: "ProductManager") -> Callable[..., "Product @pytest.fixture -def payout_config(request) -> "PayoutConfig": +def payout_config(request: Request) -> PayoutConfig: from generalresearch.models.thl.product import ( PayoutConfig, PayoutTransformation, @@ -348,8 +348,8 @@ def payout_config(request) -> "PayoutConfig": @pytest.fixture def product_user_wallet_yes( - payout_config: "PayoutConfig", product_manager: "ProductManager" -) -> "Product": + payout_config: PayoutConfig, product_manager: ProductManager +) -> Product: from generalresearch.models.thl.product import UserWalletConfig return product_manager.create_dummy( @@ -358,7 +358,7 @@ def product_user_wallet_yes( @pytest.fixture -def product_user_wallet_no(product_manager: "ProductManager") -> "Product": +def product_user_wallet_no(product_manager: ProductManager) -> Product: from generalresearch.models.thl.product import UserWalletConfig return product_manager.create_dummy( @@ -368,8 +368,8 @@ def product_user_wallet_no(product_manager: "ProductManager") -> "Product": @pytest.fixture def product_amt_true( - product_manager: "ProductManager", payout_config: "PayoutConfig" -) -> "Product": + product_manager: ProductManager, payout_config: PayoutConfig +) -> Product: from generalresearch.models.thl.product import UserWalletConfig return product_manager.create_dummy( @@ -380,17 +380,19 @@ def product_amt_true( @pytest.fixture def bp_payout_factory( - thl_lm: "ThlLedgerManager", - product_manager: "ProductManager", - business_payout_event_manager: "BusinessPayoutEventManager", -) -> Callable[..., "BrokerageProductPayoutEvent"]: + thl_lm: ThlLedgerManager, + product_manager: ProductManager, + business_payout_event_manager: BusinessPayoutEventManager, +) -> Callable[..., BrokerageProductPayoutEvent]: def _inner( - product: Optional["Product"] = None, - amount: Optional["USDCent"] = None, - ext_ref_id: Optional[str] = None, - created: Optional[AwareDatetime] = None, - ) -> "BrokerageProductPayoutEvent": + product: Product | None = None, + amount: USDCent | None = None, + ext_ref_id: str | None = None, + created: AwareDatetime | None = None, + skip_wallet_balance_check: bool = False, + skip_one_per_day_check: bool = False, + ) -> BrokerageProductPayoutEvent: from generalresearch.currency import USDCent product = product or product_manager.create_dummy() @@ -411,120 +413,49 @@ def bp_payout_factory( @pytest.fixture -def business(request, business_manager: "BusinessManager") -> "Business": +def business(request, business_manager: BusinessManager) -> Business: return business_manager.create_dummy() @pytest.fixture def business_address( - request, business: "Business", business_address_manager: "BusinessAddressManager" -) -> "BusinessAddress": + request, business: Business, business_address_manager: BusinessAddressManager +) -> BusinessAddress: return business_address_manager.create_dummy(business_id=business.id) @pytest.fixture def business_bank_account( request, - business: "Business", - business_bank_account_manager: "BusinessBankAccountManager", -) -> "BusinessBankAccount": + business: Business, + business_bank_account_manager: BusinessBankAccountManager, +) -> BusinessBankAccount: return business_bank_account_manager.create_dummy(business_id=business.id) @pytest.fixture -def team(request, team_manager: "TeamManager") -> "Team": +def team(request, team_manager: TeamManager) -> Team: return team_manager.create_dummy() @pytest.fixture -def gr_user(gr_um: "GRUserManager") -> "GRUser": - return gr_um.create_dummy() - - -@pytest.fixture -def gr_user_cache( - gr_user: "GRUser", - gr_db: PostgresConfig, - thl_web_rr: PostgresConfig, - gr_redis_config: RedisConfig, -) -> "GRUser": - gr_user.set_cache( - pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config - ) - return gr_user - - -@pytest.fixture -def gr_user_factory(gr_um: "GRUserManager") -> Callable[..., "GRUser"]: - - def _inner(): - return gr_um.create_dummy() - - return _inner - - -@pytest.fixture() -def gr_user_token( - gr_user: "GRUser", gr_tm: "GRTokenManager", gr_db: PostgresConfig -) -> "GRToken": - gr_tm.create(user_id=gr_user.id) - gr_user.prefetch_token(pg_config=gr_db) - - res = gr_user.token - assert res is not None, "GRToken should exist after creation and prefetching" - return res - - -@pytest.fixture() -def gr_user_token_header(gr_user_token: "GRToken") -> Dict[str, str]: - return gr_user_token.auth_header - - -@pytest.fixture(scope="function") -def membership( - request, team: "Team", gr_user: "GRUser", team_manager: "TeamManager" -) -> "Membership": - assert team.id, "Team must be saved" - assert gr_user.id, "GRUser must be saved" - return team_manager.add_user(team=team, gr_user=gr_user) - - -@pytest.fixture(scope="function") -def membership_factory( - team: "Team", - gr_user: "GRUser", - membership_manager: "MembershipManager", - team_manager: "TeamManager", - gr_um: "GRUserManager", -) -> Callable[..., "Membership"]: - - def _inner(**kwargs) -> "Membership": - _team = kwargs.get("team", team_manager.create_dummy()) - _gr_user = kwargs.get("gr_user", gr_um.create_dummy()) - - return membership_manager.create(team=_team, gr_user=_gr_user) - - return _inner - - -@pytest.fixture -def audit_log(audit_log_manager: "AuditLogManager", user: "User") -> "AuditLog": +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"]: + audit_log_manager: AuditLogManager, +) -> Callable[..., AuditLog]: def _inner( user_id: PositiveInt, - level: Optional["AuditLogLevel"] = None, - event_type: Optional[str] = None, - event_msg: Optional[str] = None, - event_value: Optional[float] = None, - ) -> "AuditLog": + 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, @@ -537,14 +468,14 @@ def audit_log_factory( @pytest.fixture -def ip_geoname(ip_geoname_manager: "IPGeonameManager") -> "IPGeoname": +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": + 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 ) @@ -552,10 +483,10 @@ def ip_information( @pytest.fixture def ip_information_factory( - ip_information_manager: "IPInformationManager", -) -> Callable[..., "IPInformation"]: + ip_information_manager: IPInformationManager, +) -> Callable[..., IPInformation]: - def _inner(ip: str, geoname: "IPGeoname", **kwargs) -> "IPInformation": + def _inner(ip: str, geoname: IPGeoname, **kwargs) -> IPInformation: return ip_information_manager.create_dummy( ip=ip, geoname_id=geoname.geoname_id, @@ -568,25 +499,25 @@ def ip_information_factory( @pytest.fixture def ip_record( - ip_record_manager: "IPRecordManager", ip_geoname: "IPGeoname", user: "User" -) -> "IPRecord": + 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"]: + ip_record_manager: IPRecordManager, user: User +) -> Callable[..., IPRecord]: - def _inner(user_id: PositiveInt, ip: Optional[str] = None) -> "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": +def buyer(buyer_manager: BuyerManager) -> Buyer: buyer_code = uuid4().hex buyer_manager.bulk_get_or_create(source=Source.TESTING, codes=[buyer_code]) b = Buyer( @@ -597,7 +528,7 @@ def buyer(buyer_manager: "BuyerManager") -> "Buyer": @pytest.fixture(scope="session") -def buyer_factory(buyer_manager: "BuyerManager") -> Callable[..., "Buyer"]: +def buyer_factory(buyer_manager: BuyerManager) -> Callable[..., Buyer]: def _inner() -> Buyer: return buyer_manager.bulk_get_or_create( @@ -608,7 +539,7 @@ def buyer_factory(buyer_manager: "BuyerManager") -> Callable[..., "Buyer"]: @pytest.fixture(scope="session") -def survey(survey_manager: "SurveyManager", buyer: "Buyer") -> "Survey": +def survey(survey_manager: SurveyManager, buyer: Buyer) -> Survey: s = Survey(source=Source.TESTING, survey_id=uuid4().hex, buyer_code=buyer.code) survey_manager.create_bulk([s]) return s @@ -616,10 +547,10 @@ def survey(survey_manager: "SurveyManager", buyer: "Buyer") -> "Survey": @pytest.fixture(scope="session") def survey_factory( - survey_manager: "SurveyManager", buyer_factory: Callable[..., "Buyer"] -) -> Callable[..., "Survey"]: + survey_manager: SurveyManager, buyer_factory: Callable[..., Buyer] +) -> Callable[..., Survey]: - def _inner(buyer: Optional[Buyer] = None) -> "Survey": + def _inner(buyer: Buyer | None = None) -> Survey: buyer = buyer or buyer_factory() s = Survey( source=Source.TESTING, diff --git a/test_utils/models/contest/__init__.py b/test_utils/models/contest/__init__.py new file mode 100644 index 0000000..e69de29 --- /dev/null +++ b/test_utils/models/contest/__init__.py diff --git a/test_utils/models/contest/conftest.py b/test_utils/models/contest/conftest.py new file mode 100644 index 0000000..e750076 --- /dev/null +++ b/test_utils/models/contest/conftest.py @@ -0,0 +1,292 @@ +from __future__ import annotations + +from datetime import datetime, timezone +from decimal import Decimal +from typing import Callable +from uuid import uuid4 + +import pytest +from fastapi import Request + +from generalresearch.currency import USDCent +from generalresearch.managers.thl.contest_manager import ContestManager +from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager +from generalresearch.models.thl.contest.contest import Contest +from generalresearch.models.thl.contest.leaderboard import ( + LeaderboardContestCreate, +) +from generalresearch.models.thl.contest.milestone import ( + MilestoneContestCreate, +) +from generalresearch.models.thl.contest.raffle import ( + RaffleContestCreate, +) +from generalresearch.models.thl.product import Product +from generalresearch.models.thl.user import User + +# === Miscellaneous === + +# === Managers === + +# === Models === + + +@pytest.fixture +def raffle_contest_create() -> RaffleContestCreate: + from generalresearch.models.thl.contest import ( + ContestEndCondition, + ContestPrize, + ) + from generalresearch.models.thl.contest.definitions import ( + ContestPrizeKind, + ContestType, + ) + from generalresearch.models.thl.contest.raffle import ( + ContestEntryType, + RaffleContestCreate, + ) + + # This is what we'll get from the fastapi endpoint + return RaffleContestCreate( + name="test", + contest_type=ContestType.RAFFLE, + entry_type=ContestEntryType.CASH, + prizes=[ + ContestPrize( + name="iPod 64GB White", + kind=ContestPrizeKind.PHYSICAL, + estimated_cash_value=USDCent(100), + ) + ], + end_condition=ContestEndCondition(target_entry_amount=USDCent(100)), + ) + + +@pytest.fixture +def raffle_contest_in_db( + product_user_wallet_yes: Product, + raffle_contest_create: RaffleContestCreate, + contest_manager: ContestManager, +) -> Contest: + return contest_manager.create( + product_id=product_user_wallet_yes.uuid, contest_create=raffle_contest_create + ) + + +@pytest.fixture +def raffle_contest( + product_user_wallet_yes: Product, raffle_contest_create: RaffleContestCreate +) -> Contest: + from generalresearch.models.thl.contest.io import contest_create_to_contest + + return contest_create_to_contest( + product_id=product_user_wallet_yes.uuid, contest_create=raffle_contest_create + ) + + +@pytest.fixture(scope="function") +def raffle_contest_factory( + product_user_wallet_yes: Product, + raffle_contest_create: RaffleContestCreate, + contest_manager: ContestManager, +) -> Callable[..., Contest]: + + def _inner(**kwargs): + raffle_contest_create.update(**kwargs) + return contest_manager.create( + product_id=product_user_wallet_yes.uuid, + contest_create=raffle_contest_create, + ) + + return _inner + + +@pytest.fixture +def milestone_contest_create() -> MilestoneContestCreate: + from generalresearch.models.thl.contest import ( + ContestPrize, + ) + from generalresearch.models.thl.contest.definitions import ( + ContestPrizeKind, + ContestType, + ) + from generalresearch.models.thl.contest.milestone import ( + ContestEntryTrigger, + MilestoneContestCreate, + MilestoneContestEndCondition, + ) + + # This is what we'll get from the fastapi endpoint + return MilestoneContestCreate( + name="Win a 50% bonus for 7 days and a $1 bonus after your first 3 completes!", + description="only valid for the first 5 users", + contest_type=ContestType.MILESTONE, + prizes=[ + ContestPrize( + name="50% for 7 days", + kind=ContestPrizeKind.PROMOTION, + estimated_cash_value=USDCent(0), + ), + ContestPrize( + name="$1 Bonus", + kind=ContestPrizeKind.CASH, + cash_amount=USDCent(1_00), + estimated_cash_value=USDCent(1_00), + ), + ], + end_condition=MilestoneContestEndCondition( + ends_at=datetime(year=2030, month=1, day=1, tzinfo=timezone.utc), + max_winners=5, + ), + entry_trigger=ContestEntryTrigger.TASK_COMPLETE, + target_amount=3, + ) + + +@pytest.fixture +def milestone_contest_in_db( + product_user_wallet_yes: Product, + milestone_contest_create: MilestoneContestCreate, + contest_manager: ContestManager, +) -> Contest: + return contest_manager.create( + product_id=product_user_wallet_yes.uuid, contest_create=milestone_contest_create + ) + + +@pytest.fixture +def milestone_contest( + product_user_wallet_yes: Product, + milestone_contest_create: MilestoneContestCreate, +) -> Contest: + from generalresearch.models.thl.contest.io import contest_create_to_contest + + return contest_create_to_contest( + product_id=product_user_wallet_yes.uuid, contest_create=milestone_contest_create + ) + + +@pytest.fixture(scope="function") +def milestone_contest_factory( + product_user_wallet_yes: Product, + milestone_contest_create: MilestoneContestCreate, + contest_manager: ContestManager, +) -> Callable[..., Contest]: + + def _inner(**kwargs): + milestone_contest_create.update(**kwargs) + return contest_manager.create( + product_id=product_user_wallet_yes.uuid, + contest_create=milestone_contest_create, + ) + + return _inner + + +@pytest.fixture +def leaderboard_contest_create( + product_user_wallet_yes: Product, +) -> LeaderboardContestCreate: + from generalresearch.models.thl.contest import ( + ContestPrize, + ) + from generalresearch.models.thl.contest.definitions import ( + ContestPrizeKind, + ContestType, + ) + from generalresearch.models.thl.contest.leaderboard import ( + LeaderboardContestCreate, + ) + + # This is what we'll get from the fastapi endpoint + return LeaderboardContestCreate( + name="test", + contest_type=ContestType.LEADERBOARD, + prizes=[ + ContestPrize( + name="$15 Cash", + estimated_cash_value=USDCent(15_00), + cash_amount=USDCent(15_00), + kind=ContestPrizeKind.CASH, + leaderboard_rank=1, + ), + ContestPrize( + name="$10 Cash", + estimated_cash_value=USDCent(10_00), + cash_amount=USDCent(10_00), + kind=ContestPrizeKind.CASH, + leaderboard_rank=2, + ), + ], + leaderboard_key=f"leaderboard:{product_user_wallet_yes.uuid}:us:daily:2025-01-01:complete_count", + ) + + +@pytest.fixture +def leaderboard_contest_in_db( + product_user_wallet_yes: Product, + leaderboard_contest_create: LeaderboardContestCreate, + contest_manager: ContestManager, +) -> Contest: + return contest_manager.create( + product_id=product_user_wallet_yes.uuid, + contest_create=leaderboard_contest_create, + ) + + +@pytest.fixture +def leaderboard_contest( + product_user_wallet_yes: Product, + leaderboard_contest_create: LeaderboardContestCreate, +): + from generalresearch.models.thl.contest.io import contest_create_to_contest + + return contest_create_to_contest( + product_id=product_user_wallet_yes.uuid, + contest_create=leaderboard_contest_create, + ) + + +@pytest.fixture(scope="function") +def leaderboard_contest_factory( + product_user_wallet_yes: Product, + leaderboard_contest_create: LeaderboardContestCreate, + contest_manager: ContestManager, +) -> Callable[..., Contest]: + + def _inner(**kwargs): + leaderboard_contest_create.update(**kwargs) + return contest_manager.create( + product_id=product_user_wallet_yes.uuid, + contest_create=leaderboard_contest_create, + ) + + return _inner + + +@pytest.fixture +def user_with_money( + request: Request, + user_factory: Callable[..., User], + product_user_wallet_yes: Product, + thl_lm: ThlLedgerManager, +) -> User: + + params = getattr(request, "param", {}) or {} + min_balance = int(params.get("min_balance", USDCent(1_00))) + + user: User = user_factory(product=product_user_wallet_yes) + wallet = thl_lm.get_account_or_create_user_wallet(user) + balance = thl_lm.get_account_balance(wallet) + todo = min_balance - balance + if todo > 0: + # # Put money in user's wallet + thl_lm.create_tx_user_bonus( + user=user, + ref_uuid=uuid4().hex, + description="bonus", + amount=Decimal(todo) / 100, + ) + print(f"wallet balance: {thl_lm.get_user_wallet_balance(user=user)}") + + return user diff --git a/test_utils/models/gr/__init__.py b/test_utils/models/gr/__init__.py new file mode 100644 index 0000000..e69de29 --- /dev/null +++ b/test_utils/models/gr/__init__.py diff --git a/test_utils/models/gr/conftest.py b/test_utils/models/gr/conftest.py new file mode 100644 index 0000000..df97306 --- /dev/null +++ b/test_utils/models/gr/conftest.py @@ -0,0 +1,213 @@ +from __future__ import annotations + +from typing import Callable +from uuid import uuid4 + +import pytest +from pydantic import PositiveInt +from pydantic_extra_types.phone_numbers import PhoneNumber + +from generalresearch.managers.gr.authentication import GRTokenManager, GRUserManager +from generalresearch.managers.gr.business import ( + BusinessAddressManager, + BusinessBankAccountManager, + BusinessManager, +) +from generalresearch.managers.gr.team import MembershipManager, TeamManager +from generalresearch.models.custom_types import UUIDStr +from generalresearch.models.gr.authentication import GRToken, GRUser +from generalresearch.models.gr.business import ( + Business, + BusinessAddress, + BusinessBankAccount, + BusinessType, + TransferMethod, +) +from generalresearch.models.gr.team import Membership, Team +from generalresearch.pg_helper import PostgresConfig +from generalresearch.redis_helper import RedisConfig + +# --- Static --- + + +# --- Factory / Database --- + + +@pytest.fixture +def gr_user_factory(gr_user_manager: GRUserManager) -> Callable[..., GRUser]: + + def _inner( + sub: str | None = None, + is_superuser: bool = False, + ) -> GRUser: + sub = sub or f"{uuid4().hex}-{uuid4().hex}" + + return gr_user_manager.create( + sub=sub, + is_superuser=is_superuser, + ) + + return _inner + + +@pytest.fixture +def gr_user_cache( + gr_user: GRUser, + gr_db: PostgresConfig, + thl_web_rr: PostgresConfig, + gr_redis_config: RedisConfig, +) -> GRUser: + gr_user.set_cache( + pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config + ) + return gr_user + + +@pytest.fixture +def gr_business_bank_account_factory( + gr_bbam: BusinessBankAccountManager, +) -> Callable[..., BusinessBankAccount]: + + def _inner( + business_id: PositiveInt, + uuid: UUIDStr | None = None, + transfer_method: TransferMethod | None = None, + account_number: str | None = None, + routing_number: str | None = None, + iban: str | None = None, + swift: str | None = None, + ): + from generalresearch.models.gr.business import TransferMethod + + return gr_bbam.create( + business_id=business_id, + uuid=uuid or uuid4().hex, + transfer_method=transfer_method or TransferMethod.ACH, + account_number=account_number or uuid4().hex[:6], + routing_number=routing_number or uuid4().hex[:6], + iban=iban or uuid4().hex[:6], + swift=swift or uuid4().hex[:6], + ) + + return _inner + + +@pytest.fixture +def gr_business_address_factory( + gr_bam: BusinessAddressManager, +) -> Callable[..., BusinessAddress]: + + def _inner( + business_id: PositiveInt, + uuid: UUIDStr | None = None, + line_1: str | None = None, + line_2: str | None = None, + city: str | None = None, + state: str | None = None, + postal_code: str | None = None, + phone_number: PhoneNumber | None = None, + country: str | None = None, + ): + uuid = uuid or uuid4().hex + line_1 = line_1 or "abc" + line_2 = line_2 or "bczx" + city = city or "Downingtown" + state = state or "CA" + postal_code = postal_code or "94041" + phone_number = None + country = country or "US" + + return gr_bam.create( + business_id=business_id, + uuid=uuid, + line_1=line_1, + line_2=line_2, + city=city, + state=state, + postal_code=postal_code, + phone_number=phone_number, + country=country, + ) + + return _inner + + +@pytest.fixture +def gr_business_factory( + gr_bm: BusinessManager, +) -> Callable[..., Business]: + + def _inner( + uuid: UUIDStr | None = None, + name: str | None = None, + team: Team | None = None, + kind: BusinessType | None = None, + tax_number: str | None = None, + ) -> Business: + from random import randint + + uuid = uuid or uuid4().hex + name = name or "< Unknown >" + tax_number = tax_number or str(randint(1, 999_999_999)) + + return gr_bm.create( + uuid=uuid, name=name, team=team, kind=kind, tax_number=tax_number + ) + + return _inner + + +@pytest.fixture +def gr_team( + gr_tm: TeamManager, +) -> Callable[..., Team]: + + def _inner(uuid: UUIDStr | None = None, name: str | None = None) -> Team: + uuid = uuid or uuid4().hex + name = name or f"name-{uuid4().hex[:12]}" + + return gr_tm.create(uuid=uuid, name=name) + + return _inner + + +@pytest.fixture() +def gr_user_token( + gr_user: GRUser, gr_tm: GRTokenManager, gr_db: PostgresConfig +) -> GRToken: + gr_tm.create(user_id=gr_user.id) + gr_user.prefetch_token(pg_config=gr_db) + + res = gr_user.token + assert res is not None, "GRToken should exist after creation and prefetching" + return res + + +@pytest.fixture() +def gr_user_token_header(gr_user_token: GRToken) -> dict[str, str]: + return gr_user_token.auth_header + + +@pytest.fixture(scope="function") +def membership(team: Team, gr_user: GRUser, team_manager: TeamManager) -> Membership: + assert team.id, "Team must be saved" + assert gr_user.id, "GRUser must be saved" + return team_manager.add_user(team=team, gr_user=gr_user) + + +@pytest.fixture(scope="function") +def membership_factory( + team: Team, + gr_user: GRUser, + membership_manager: MembershipManager, + team_manager: TeamManager, + gr_um: GRUserManager, +) -> Callable[..., Membership]: + + def _inner(**kwargs) -> Membership: + _team = kwargs.get("team", team_manager.create_dummy()) + _gr_user = kwargs.get("gr_user", gr_um.create_dummy()) + + return membership_manager.create(team=_team, gr_user=_gr_user) + + return _inner diff --git a/test_utils/models/ledger/__init__.py b/test_utils/models/ledger/__init__.py new file mode 100644 index 0000000..e69de29 --- /dev/null +++ b/test_utils/models/ledger/__init__.py diff --git a/test_utils/models/ledger/conftest.py b/test_utils/models/ledger/conftest.py new file mode 100644 index 0000000..5bef113 --- /dev/null +++ b/test_utils/models/ledger/conftest.py @@ -0,0 +1,724 @@ +from __future__ import annotations + +from datetime import datetime +from decimal import Decimal +from random import randint +from typing import TYPE_CHECKING, Callable +from uuid import uuid4 + +import pytest +from fastapi import Request + +from generalresearch.currency import USDCent +from generalresearch.managers.base import PostgresManager +from test_utils.models.conftest import ( + payout_config, + product_amt_true, + product_user_wallet_no, + product_user_wallet_yes, + session, + session_factory, + user_factory, + wall, + wall_factory, +) + +_ = ( + user_factory, + product_user_wallet_no, + wall, + product_amt_true, + product_user_wallet_yes, + session_factory, + session, + wall_factory, + payout_config, +) + +if TYPE_CHECKING: + + from generalresearch.currency import LedgerCurrency + from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager + from generalresearch.managers.thl.ledger_manager.thl_ledger import ( + ThlLedgerManager, + ) + from generalresearch.managers.thl.payout import ( + BrokerageProductPayoutEventManager, + BusinessPayoutEventManager, + ) + from generalresearch.managers.thl.session import SessionManager + from generalresearch.managers.thl.wall import WallManager + from generalresearch.models.thl.ledger import ( + LedgerAccount, + LedgerTransaction, + ) + from generalresearch.models.thl.payout import ( + BrokerageProductPayoutEvent, + ) + from generalresearch.models.thl.product import Product + from generalresearch.models.thl.session import Session + from generalresearch.models.thl.user import User + + +@pytest.fixture +def ledger_account( + request: Request, lm: LedgerManager, currency: LedgerCurrency +) -> LedgerAccount: + from generalresearch.models.thl.ledger import ( + AccountType, + Direction, + LedgerAccount, + ) + + account_type = getattr(request, "account_type", AccountType.CASH) + direction = getattr(request, "direction", Direction.CREDIT) + + acct_uuid = uuid4().hex + qn = f"{currency}:{account_type}:{acct_uuid}" + + acct_model = LedgerAccount( + uuid=acct_uuid, + display_name=f"test-{acct_uuid}", + currency=currency, + qualified_name=qn, + account_type=account_type, + normal_balance=direction, + ) + return lm.create_account(account=acct_model) + + +@pytest.fixture +def ledger_account_factory( + request: Request, + thl_lm: ThlLedgerManager, + lm: LedgerManager, + currency: LedgerCurrency, +) -> Callable[..., LedgerAccount]: + + from generalresearch.models.thl.ledger import ( + AccountType, + Direction, + LedgerAccount, + ) + + def _inner( + product: Product, + account_type: AccountType = AccountType.CASH, + direction: Direction = Direction.CREDIT, + ) -> LedgerAccount: + thl_lm.get_account_or_create_bp_wallet(product=product) + acct_uuid = uuid4().hex + qn = f"{currency}:{account_type}:{acct_uuid}" + + acct_model = LedgerAccount( + uuid=acct_uuid, + display_name=f"test-{acct_uuid}", + currency=currency, + qualified_name=qn, + account_type=account_type, + normal_balance=direction, + ) + return lm.create_account(account=acct_model) + + return _inner + + +@pytest.fixture +def ledger_account_credit( + request: Request, lm: LedgerManager, currency: LedgerCurrency +) -> LedgerAccount: + from generalresearch.models.thl.ledger import AccountType, Direction + + account_type = AccountType.REVENUE + acct_uuid = uuid4().hex + + qn = f"{currency}:{account_type}:{acct_uuid}" + from generalresearch.models.thl.ledger import LedgerAccount + + acct_model = LedgerAccount( + uuid=acct_uuid, + display_name=f"test-{acct_uuid}", + currency=currency, + qualified_name=qn, + account_type=account_type, + normal_balance=Direction.CREDIT, + ) + return lm.create_account(account=acct_model) + + +@pytest.fixture +def ledger_account_debit( + request: Request, lm: LedgerManager, currency: LedgerCurrency +) -> LedgerAccount: + from generalresearch.models.thl.ledger import AccountType, Direction + + account_type = AccountType.EXPENSE + acct_uuid = uuid4().hex + + qn = f"{currency}:{account_type}:{acct_uuid}" + from generalresearch.models.thl.ledger import LedgerAccount + + acct_model = LedgerAccount( + uuid=acct_uuid, + display_name=f"test-{acct_uuid}", + currency=currency, + qualified_name=qn, + account_type=account_type, + normal_balance=Direction.DEBIT, + ) + return lm.create_account(account=acct_model) + + +@pytest.fixture +def tag(request: Request, lm: LedgerManager) -> str: + from generalresearch.currency import LedgerCurrency + + return ( + request.param + if hasattr(request, "tag") + else f"{LedgerCurrency.TEST}:{uuid4().hex}" + ) + + +@pytest.fixture +def usd_cent(request: Request) -> USDCent: + amount = randint(99, 9_999) + return request.param if hasattr(request, "usd_cent") else USDCent(amount) + + +@pytest.fixture +def bp_payout_event( + product: Product, + usd_cent: USDCent, + business_payout_event_manager: BusinessPayoutEventManager, + thl_lm: ThlLedgerManager, +) -> BrokerageProductPayoutEvent: + + return business_payout_event_manager.create_bp_payout_event( + thl_ledger_manager=thl_lm, + product=product, + amount=usd_cent, + skip_wallet_balance_check=True, + skip_one_per_day_check=True, + ) + + +@pytest.fixture +def bp_payout_event_factory( + brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager, + thl_lm: ThlLedgerManager, +) -> Callable[..., BrokerageProductPayoutEvent]: + + def _inner( + product: Product, usd_cent: USDCent, ext_ref_id: str | None = None + ) -> BrokerageProductPayoutEvent: + + return brokerage_product_payout_event_manager.create_bp_payout_event( + thl_ledger_manager=thl_lm, + product=product, + amount=usd_cent, + ext_ref_id=ext_ref_id, + skip_wallet_balance_check=True, + skip_one_per_day_check=True, + ) + + return _inner + + +@pytest.fixture +def currency(lm: LedgerManager) -> LedgerCurrency: + # return request.param if hasattr(request, "currency") else LedgerCurrency.TEST + assert lm.currency, "LedgerManager must have a currency specified for these tests" + return lm.currency + + +@pytest.fixture +def tx_metadata(request: Request) -> dict[str, str] | None: + return ( + request.param + if hasattr(request, "tx_metadata") + else {f"key-{uuid4().hex[:10]}": uuid4().hex} + ) + + +@pytest.fixture +def ledger_tx( + request: Request, + ledger_account_credit: LedgerAccount, + ledger_account_debit: LedgerAccount, + tag: str, + currency: LedgerCurrency, + tx_metadata: dict[str, str] | None, + lm: LedgerManager, +) -> LedgerTransaction: + from generalresearch.models.thl.ledger import Direction, LedgerEntry + + amount = int(Decimal("1.00") * 100) + + entries = [ + LedgerEntry( + direction=Direction.CREDIT, + account_uuid=ledger_account_credit.uuid, + amount=amount, + ), + LedgerEntry( + direction=Direction.DEBIT, + account_uuid=ledger_account_debit.uuid, + amount=amount, + ), + ] + + return lm.create_tx(entries=entries, tag=tag, metadata=tx_metadata) + + +@pytest.fixture +def create_main_accounts( + lm: LedgerManager, currency: LedgerCurrency +) -> Callable[..., None]: + + def _inner() -> None: + from generalresearch.models.thl.ledger import ( + AccountType, + Direction, + LedgerAccount, + ) + + account = LedgerAccount( + display_name="Cash flow task complete", + qualified_name=f"{currency.value}:revenue:task_complete", + normal_balance=Direction.CREDIT, + account_type=AccountType.REVENUE, + currency=lm.currency, + ) + lm.get_account_or_create(account=account) + + account = LedgerAccount( + display_name="Operating Cash Account", + qualified_name=f"{currency.value}:cash", + normal_balance=Direction.DEBIT, + account_type=AccountType.CASH, + currency=currency, + ) + + lm.get_account_or_create(account=account) + + return _inner + + +@pytest.fixture +def delete_ledger_db(thl_web_rw: PostgresManager) -> Callable[..., None]: + + def _inner(): + for table in [ + "ledger_transactionmetadata", + "ledger_entry", + "ledger_transaction", + "ledger_account", + ]: + thl_web_rw.execute_write( + query=f"DELETE FROM {table};", + ) + + return _inner + + +@pytest.fixture +def wipe_main_accounts( + thl_web_rw: PostgresManager, lm: LedgerManager, currency: LedgerCurrency +) -> Callable[..., None]: + + def _inner() -> None: + db_table = thl_web_rw.db_name + qual_names = [ + f"{currency.value}:revenue:task_complete", + f"{currency.value}:cash", + ] + + res = thl_web_rw.execute_sql_query( + query=f""" + SELECT lt.id as ltid, le.id as leid, tmd.id as tmdid, la.uuid as lauuid + FROM `{db_table}`.`ledger_transaction` AS lt + LEFT JOIN `{db_table}`.ledger_entry le + ON lt.id = le.transaction_id + LEFT JOIN `{db_table}`.ledger_account la + ON la.uuid = le.account_id + LEFT JOIN `{db_table}`.ledger_transactionmetadata tmd + ON lt.id = tmd.transaction_id + WHERE la.qualified_name IN %s + """, + params=[qual_names], + ) + + lt = {x["ltid"] for x in res if x["ltid"]} + le = {x["leid"] for x in res if x["leid"]} + tmd = {x["tmdid"] for x in res if x["tmdid"]} + la = {x["lauuid"] for x in res if x["lauuid"]} + + thl_web_rw.execute_sql_query( + query=f""" + DELETE FROM `{db_table}`.`ledger_transactionmetadata` + WHERE id IN %s + """, + params=[tmd], + commit=True, + ) + + thl_web_rw.execute_sql_query( + query=f""" + DELETE FROM `{db_table}`.`ledger_entry` + WHERE id IN %s + """, + params=[le], + commit=True, + ) + + thl_web_rw.execute_sql_query( + query=f""" + DELETE FROM `{db_table}`.`ledger_transaction` + WHERE id IN %s + """, + params=[lt], + commit=True, + ) + + thl_web_rw.execute_sql_query( + query=f""" + DELETE FROM `{db_table}`.`ledger_account` + WHERE uuid IN %s + """, + params=[la], + commit=True, + ) + + return _inner + + +@pytest.fixture +def account_cash(lm: LedgerManager, currency: LedgerCurrency) -> LedgerAccount: + from generalresearch.models.thl.ledger import ( + AccountType, + Direction, + LedgerAccount, + ) + + account = LedgerAccount( + display_name="Operating Cash Account", + qualified_name=f"{currency.value}:cash", + normal_balance=Direction.DEBIT, + account_type=AccountType.CASH, + currency=currency, + ) + return lm.get_account_or_create(account=account) + + +@pytest.fixture +def account_revenue_task_complete( + lm: LedgerManager, currency: LedgerCurrency +) -> LedgerAccount: + from generalresearch.models.thl.ledger import ( + AccountType, + Direction, + LedgerAccount, + ) + + account = LedgerAccount( + display_name="Cash flow task complete", + qualified_name=f"{currency.value}:revenue:task_complete", + normal_balance=Direction.CREDIT, + account_type=AccountType.REVENUE, + currency=currency, + ) + return lm.get_account_or_create(account=account) + + +@pytest.fixture +def account_expense_tango(lm: LedgerManager, currency: LedgerCurrency) -> LedgerAccount: + from generalresearch.models.thl.ledger import ( + AccountType, + Direction, + LedgerAccount, + ) + + account = LedgerAccount( + display_name="Tango Fee", + qualified_name=f"{currency.value}:expense:tango_fee", + normal_balance=Direction.DEBIT, + account_type=AccountType.EXPENSE, + currency=currency, + ) + return lm.get_account_or_create(account=account) + + +@pytest.fixture +def user_account_user_wallet( + lm: LedgerManager, user: User, currency: LedgerCurrency +) -> LedgerAccount: + from generalresearch.models.thl.ledger import ( + AccountType, + Direction, + LedgerAccount, + ) + + account = LedgerAccount( + display_name=f"{user.uuid} Wallet", + qualified_name=f"{currency.value}:user_wallet:{user.uuid}", + normal_balance=Direction.CREDIT, + account_type=AccountType.USER_WALLET, + reference_type="user", + reference_uuid=user.uuid, + currency=currency, + ) + return lm.get_account_or_create(account=account) + + +@pytest.fixture +def product_account_bp_wallet( + lm: LedgerManager, product: Product, currency: LedgerCurrency +) -> LedgerAccount: + from generalresearch.models.thl.ledger import ( + AccountType, + Direction, + LedgerAccount, + ) + + account = LedgerAccount.model_validate( + { + "display_name": f"{product.name} Wallet", + "qualified_name": f"{currency.value}:bp_wallet:{product.uuid}", + "normal_balance": Direction.CREDIT, + "account_type": AccountType.BP_WALLET, + "reference_type": "bp", + "reference_uuid": product.uuid, + "currency": currency, + } + ) + return lm.get_account_or_create(account=account) + + +@pytest.fixture +def setup_accounts( + product_factory: Callable[..., Product], + lm: LedgerManager, + user: User, + currency: LedgerCurrency, +) -> None: + from generalresearch.models.thl.ledger import ( + AccountType, + Direction, + LedgerAccount, + ) + + # BP's wallet and a revenue from their commissions account. + p1 = product_factory() + + account = LedgerAccount( + display_name=f"Revenue from {p1.name} commission", + qualified_name=f"{currency.value}:revenue:bp_commission:{p1.uuid}", + normal_balance=Direction.CREDIT, + account_type=AccountType.REVENUE, + reference_type="bp", + reference_uuid=p1.uuid, + currency=currency, + ) + lm.get_account_or_create(account=account) + + account = LedgerAccount.model_validate( + { + "display_name": f"{p1.name} Wallet", + "qualified_name": f"{currency.value}:bp_wallet:{p1.uuid}", + "normal_balance": Direction.CREDIT, + "account_type": AccountType.BP_WALLET, + "reference_type": "bp", + "reference_uuid": p1.uuid, + "currency": currency, + } + ) + lm.get_account_or_create(account=account) + + # BP's wallet, user's wallet, and a revenue from their commissions account. + p2 = product_factory() + account = LedgerAccount( + display_name=f"Revenue from {p2.name} commission", + qualified_name=f"{currency.value}:revenue:bp_commission:{p2.uuid}", + normal_balance=Direction.CREDIT, + account_type=AccountType.REVENUE, + reference_type="bp", + reference_uuid=p2.uuid, + currency=currency, + ) + lm.get_account_or_create(account) + + account = LedgerAccount( + display_name=f"{p2.name} Wallet", + qualified_name=f"{currency.value}:bp_wallet:{p2.uuid}", + normal_balance=Direction.CREDIT, + account_type=AccountType.BP_WALLET, + reference_type="bp", + reference_uuid=p2.uuid, + currency=currency, + ) + lm.get_account_or_create(account) + + account = LedgerAccount( + display_name=f"{user.uuid} Wallet", + qualified_name=f"{currency.value}:user_wallet:{user.uuid}", + normal_balance=Direction.CREDIT, + account_type=AccountType.USER_WALLET, + reference_type="user", + reference_uuid=user.uuid, + currency="test", + ) + lm.get_account_or_create(account=account) + + +@pytest.fixture +def session_with_tx_factory( + session_factory: Callable[..., Session], + session_manager: SessionManager, + wall_manager: WallManager, + utc_hour_ago: datetime, + thl_lm: ThlLedgerManager, +) -> Callable[..., Session]: + + from generalresearch.models.thl.session import ( + Status, + StatusCode1, + ) + + def _inner( + user: User, + final_status: Status = Status.COMPLETE, + wall_req_cpi: Decimal = Decimal(".50"), + started: datetime = utc_hour_ago, + ) -> Session: + s: Session = session_factory( + user=user, + wall_count=2, + final_status=final_status, + wall_req_cpi=wall_req_cpi, + started=started, + ) + last_wall = s.wall_events[-1] + + wall_manager.finish( + wall=last_wall, + status=Status.COMPLETE, + status_code_1=StatusCode1.COMPLETE, + finished=last_wall.finished, + ) + + status, status_code_1 = s.determine_session_status() + _, _, bp_pay, user_pay = s.determine_payments() + session_manager.finish_with_status( + session=s, + finished=last_wall.finished, + payout=bp_pay, + user_payout=user_pay, + status=status, + status_code_1=status_code_1, + ) + + thl_lm.create_tx_task_complete( + wall=last_wall, + user=user, + created=last_wall.finished, + force=True, + ) + + thl_lm.create_tx_bp_payment(session=s, created=last_wall.finished, force=True) + + return s + + return _inner + + +@pytest.fixture +def adj_to_fail_with_tx_factory( + session_manager: SessionManager, + wall_manager: WallManager, + thl_lm: ThlLedgerManager, +) -> Callable[..., None]: + from datetime import timedelta + + from generalresearch.models.thl.definitions import WallAdjustedStatus + + def _inner( + session: Session, + created: datetime, + ) -> None: + w1 = wall_manager.get_wall_events(session_id=session.id)[-1] + + # This is defined in `thl-grpc/thl/user_quality_history/recons.py:150` + # so we can't use it as part of this test anyway to add rows to the + # thl_taskadjustment table anyway.. until we created a + # TaskAdjustment Manager to put into generalresearch! + + # create_task_adjustment_event( + # wall, + # user, + # adjusted_status, + # amount_usd=amount_usd, + # alert_time=alert_time, + # ext_status_code=ext_status_code, + # ) + + wall_manager.adjust_status( + wall=w1, + adjusted_status=WallAdjustedStatus.ADJUSTED_TO_FAIL, + adjusted_cpi=Decimal("0.00"), + adjusted_timestamp=created, + ) + + thl_lm.create_tx_task_adjustment( + wall=w1, + user=session.user, + created=created + timedelta(milliseconds=1), + ) + + session.wall_events = wall_manager.get_wall_events(session_id=session.id) + session_manager.adjust_status(session=session) + + thl_lm.create_tx_bp_adjustment( + session=session, created=created + timedelta(milliseconds=2) + ) + + return _inner + + +@pytest.fixture +def adj_to_complete_with_tx_factory( + session_manager: SessionManager, + wall_manager: WallManager, + thl_lm: ThlLedgerManager, +) -> Callable[..., None]: + from datetime import timedelta + + from generalresearch.models.thl.definitions import WallAdjustedStatus + + def _inner( + session: Session, + created: datetime, + ) -> None: + w1 = wall_manager.get_wall_events(session_id=session.id)[-1] + + wall_manager.adjust_status( + wall=w1, + adjusted_status=WallAdjustedStatus.ADJUSTED_TO_COMPLETE, + adjusted_cpi=w1.req_cpi, + adjusted_timestamp=created, + ) + + thl_lm.create_tx_task_adjustment( + wall=w1, + user=session.user, + created=created + timedelta(milliseconds=1), + ) + + session.wall_events = wall_manager.get_wall_events(session_id=session.id) + session_manager.adjust_status(session=session) + + thl_lm.create_tx_bp_adjustment( + session=session, created=created + timedelta(milliseconds=2) + ) + + return _inner diff --git a/test_utils/models/network/__init__.py b/test_utils/models/network/__init__.py new file mode 100644 index 0000000..e69de29 --- /dev/null +++ b/test_utils/models/network/__init__.py diff --git a/test_utils/models/network/conftest.py b/test_utils/models/network/conftest.py new file mode 100644 index 0000000..abfbc18 --- /dev/null +++ b/test_utils/models/network/conftest.py @@ -0,0 +1,144 @@ +import os +from datetime import datetime, timedelta, timezone +from uuid import uuid4 + +import pytest +from fastapi import Request + +from generalresearch.managers.network.label import IPLabelManager +from generalresearch.managers.network.tool_run import ToolRunManager +from generalresearch.models.network.definitions import IPProtocol +from generalresearch.models.network.mtr.parser import parse_mtr_output +from generalresearch.models.network.mtr.result import MTRResult +from generalresearch.models.network.nmap.parser import parse_nmap_xml +from generalresearch.models.network.nmap.result import NmapResult +from generalresearch.models.network.rdns.parser import parse_rdns_output +from generalresearch.models.network.rdns.result import RDNSResult +from generalresearch.models.network.tool_run import MTRRun, NmapRun, RDNSRun, Status +from generalresearch.models.network.tool_run_command import ( + MTRRunCommand, + MTRRunCommandOptions, + NmapRunCommand, + NmapRunCommandOptions, + RDNSRunCommand, + RDNSRunCommandOptions, +) +from generalresearch.pg_helper import PostgresConfig + + +@pytest.fixture(scope="session") +def scan_group_id() -> str: + return uuid4().hex + + +@pytest.fixture(scope="session") +def iplabel_manager(thl_web_rw: PostgresConfig) -> IPLabelManager: + return IPLabelManager(pg_config=thl_web_rw) + + +@pytest.fixture(scope="session") +def toolrun_manager(thl_web_rw: PostgresConfig) -> ToolRunManager: + return ToolRunManager(pg_config=thl_web_rw) + + +@pytest.fixture(scope="session") +def nmap_raw_output(request: Request) -> str: + fp = os.path.join(request.config.rootpath, "data/nmaprun1.xml") + with open(fp) as f: + data = f.read() + return data + + +@pytest.fixture(scope="session") +def nmap_result(nmap_raw_output: str) -> NmapResult: + return parse_nmap_xml(nmap_raw_output) + + +@pytest.fixture(scope="session") +def nmap_run(nmap_result: NmapResult, scan_group_id: str): + r = nmap_result + config = NmapRunCommand( + command="nmap", + options=NmapRunCommandOptions( + ip=r.target_ip, ports="22-1000,11000,1100,3389,61232", top_ports=None + ), + ) + return NmapRun( + tool_version=r.version, + status=Status.SUCCESS, + ip=r.target_ip, + started_at=r.started_at, + finished_at=r.finished_at, + raw_command=config.to_command_str(), + scan_group_id=scan_group_id, + config=config, + parsed=r, + ) + + +@pytest.fixture(scope="session") +def dig_raw_output() -> str: + return "156.32.33.45.in-addr.arpa. 300 IN PTR scanme.nmap.org." + + +@pytest.fixture(scope="session") +def rdns_result(dig_raw_output: str) -> RDNSResult: + return parse_rdns_output(ip="45.33.32.156", raw=dig_raw_output) + + +@pytest.fixture(scope="session") +def rdns_run(rdns_result: RDNSResult, scan_group_id: str): + r = rdns_result + ip = "45.33.32.156" + utc_now = datetime.now(tz=timezone.utc) + config = RDNSRunCommand(command="dig", options=RDNSRunCommandOptions(ip=ip)) + return RDNSRun( + tool_version="1.2.3", + status=Status.SUCCESS, + ip=ip, + started_at=utc_now, + finished_at=utc_now + timedelta(seconds=1), + raw_command=config.to_command_str(), + scan_group_id=scan_group_id, + config=config, + parsed=r, + ) + + +@pytest.fixture(scope="session") +def mtr_raw_output(request: Request) -> str: + fp = os.path.join(request.config.rootpath, "data/mtr_fatbeam.json") + with open(fp) as f: + data = f.read() + return data + + +@pytest.fixture(scope="session") +def mtr_result(mtr_raw_output: str) -> MTRResult: + return parse_mtr_output(mtr_raw_output, port=443, protocol=IPProtocol.TCP) + + +@pytest.fixture(scope="session") +def mtr_run(mtr_result: MTRResult, scan_group_id: str): + r = mtr_result + utc_now = datetime.now(tz=timezone.utc) + config = MTRRunCommand( + command="mtr", + options=MTRRunCommandOptions( + ip=r.destination, protocol=IPProtocol.TCP, port=443 + ), + ) + + return MTRRun( + tool_version="1.2.3", + status=Status.SUCCESS, + ip=r.destination, + started_at=utc_now, + finished_at=utc_now + timedelta(seconds=1), + raw_command=config.to_command_str(), + scan_group_id=scan_group_id, + config=config, + parsed=r, + facility_id=1, + source_ip="1.2.3.4", + ) diff --git a/test_utils/models/thl/__init__.py b/test_utils/models/thl/__init__.py new file mode 100644 index 0000000..e69de29 --- /dev/null +++ b/test_utils/models/thl/__init__.py diff --git a/test_utils/models/thl/conftest.py b/test_utils/models/thl/conftest.py new file mode 100644 index 0000000..cf8d2fa --- /dev/null +++ b/test_utils/models/thl/conftest.py @@ -0,0 +1,434 @@ +from __future__ import annotations + +from datetime import datetime, timezone +from decimal import ROUND_DOWN, Decimal +from random import choice as rand_choice +from random import choice as rchoice +from random import randint, random +from typing import Any, Callable +from uuid import uuid4 + +import faker +import pytest +from pydantic import PositiveInt + +from generalresearch.managers.thl.ipinfo import IPGeonameManager, IPInformationManager +from generalresearch.managers.thl.payout import UserPayoutEventManager +from generalresearch.managers.thl.product import ProductManager +from generalresearch.managers.thl.session import SessionManager +from generalresearch.managers.thl.user_manager.user_manager import UserManager +from generalresearch.managers.thl.userhealth import AuditLogManager, IPRecordManager +from generalresearch.managers.thl.wall import WallManager +from generalresearch.models import DeviceType +from generalresearch.models.custom_types import ( + AwareDatetimeISO, + IPvAnyAddressStr, + UUIDStr, +) +from generalresearch.models.legacy.bucket import Bucket +from generalresearch.models.thl.definitions import ( + PayoutStatus, +) +from generalresearch.models.thl.ipinfo import IPGeoname, IPInformation, UserType +from generalresearch.models.thl.payout import UserPayoutEvent +from generalresearch.models.thl.product import ( + PayoutConfig, + Product, + ProfilingConfig, + SessionConfig, + SourcesConfig, + SupplyConfig, + UserCreateConfig, + UserHealthConfig, + UserWalletConfig, +) +from generalresearch.models.thl.session import ( + Session, + Source, + Status, + 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.models.thl.wallet import PayoutType +from generalresearch.models.thl.wallet.cashout_method import CashMailOrderData + +fake = faker.Faker() + + +@pytest.fixture +def wall_status() -> Status: + return Status.COMPLETE + + +@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]: + + def _inner( + session_id: int | None = None, + user_id: int | None = None, + started: datetime | None = None, + source: Source | None = None, + req_survey_id: str | None = None, + req_cpi: Decimal | None = None, + buyer_id: str | None = None, + uuid_id: str | None = None, + ): + """To be used in tests, where we don't care about certain fields""" + + user_id = user_id or fake.random_int(min=1, max=2_147_483_648) + started = started or fake.date_time_between( + start_date=datetime(year=1900, month=1, day=1, tzinfo=timezone.utc), + end_date=datetime.now(tz=timezone.utc), + tzinfo=timezone.utc, + ) + + if session_id is None: + # session = SessionManager(pg_config=self.pg_config).create_dummy( + # started=started + # ) + session = session_factory() + session_id = session.id + + source = source or rchoice(list(Source)) + req_survey_id = req_survey_id or uuid4().hex + req_cpi = req_cpi or Decimal(fake.random_int(min=1, max=150) / 100).quantize( + Decimal(".01"), rounding=ROUND_DOWN + ) + + return wall_manager.create( + session_id=session_id, + user_id=user_id, + started=started, + source=source, + req_survey_id=req_survey_id, + req_cpi=req_cpi, + buyer_id=buyer_id, + uuid_id=uuid_id, + ) + + return _inner + + +@pytest.fixture +def product_factory(product_manager: ProductManager) -> Callable[..., Product]: + + def _inner( + product_id: UUIDStr | None = None, + team_id: UUIDStr | None = None, + business_id: UUIDStr | None = None, + name: str | None = None, + redirect_url: str | None = None, + harmonizer_domain: str | None = None, + commission_pct: Decimal = Decimal("0.05000"), + sources_config: SourcesConfig | SupplyConfig | None = None, + payout_config: PayoutConfig | None = None, + session_config: SessionConfig | None = None, + profiling_config: ProfilingConfig | None = None, + user_wallet_config: UserWalletConfig | None = None, + user_create_config: UserCreateConfig | None = None, + user_health_config: UserHealthConfig | None = None, + ) -> Product: + """To be used in tests, where we don't care about certain fields""" + product_id = product_id if product_id else uuid4().hex + team_id = team_id if team_id else uuid4().hex + name = name if name else f"name-{product_id[:12]}" + redirect_url = redirect_url if redirect_url else "https://www.example.com/" + + return product_manager.create( + product_id=product_id, + team_id=team_id, + business_id=business_id, + name=name, + redirect_url=redirect_url, + harmonizer_domain=harmonizer_domain, + commission_pct=commission_pct, + sources_config=sources_config, + payout_config=payout_config, + session_config=session_config, + profiling_config=profiling_config, + user_wallet_config=user_wallet_config, + user_create_config=user_create_config, + user_health_config=user_health_config, + ) + + return _inner + + +@pytest.fixture +def session_factory(session_manager: SessionManager): + + def _inner( + # -- Create Dummy "optional" -- # + started: datetime | None = None, + user: User | None = None, + # -- Optional -- # + country_iso: str | None = None, + device_type: DeviceType | None = None, + ip: str | None = None, + bucket: Bucket | None = None, + url_metadata: dict[str, str] | None = None, + uuid_id: str | None = None, + ) -> Session: + """To be used in tests, where we don't care about certain fields""" + started = started or fake.date_time_between( + start_date=datetime(year=1900, month=1, day=1, tzinfo=timezone.utc), + end_date=datetime(year=2000, month=1, day=1, tzinfo=timezone.utc), + tzinfo=timezone.utc, + ) + user = user or User( + user_id=fake.random_int(min=1, max=2_147_483_648), uuid=uuid4().hex + ) + + return session_manager.create( + started=started, + user=user, + country_iso=country_iso, + device_type=device_type, + ip=ip, + bucket=bucket, + url_metadata=url_metadata, + uuid_id=uuid_id, + ) + + return _inner + + +@pytest.fixture +def ipgeoname_factory(ipgeoname_manager: IPGeonameManager) -> Callable[..., IPGeoname]: + + def _inner( + geoname_id: PositiveInt | None = None, + continent_code: str | None = None, + continent_name: str | None = None, + country_iso: str | None = None, + country_name: str | None = None, + subdivision_1_iso: str | None = None, + subdivision_1_name: str | None = None, + subdivision_2_iso: str | None = None, + subdivision_2_name: str | None = None, + city_name: str | None = None, + metro_code: int | None = None, + 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, + ) + + return _inner + + +def ipinformation_factory( + ipinformation_manager: IPInformationManager, +) -> Callable[..., IPInformation]: + + def _inner( + ip: IPvAnyAddressStr | None = None, + geoname_id: PositiveInt | None = None, + country_iso: str | None = None, + registered_country_iso: str | None = None, + is_anonymous: bool | None = None, + is_anonymous_vpn: bool | None = None, + is_hosting_provider: bool | None = None, + is_public_proxy: bool | None = None, + is_tor_exit_node: bool | None = None, + is_residential_proxy: bool | None = None, + autonomous_system_number: PositiveInt | None = None, + autonomous_system_organization: str | None = None, + domain: str | None = None, + isp: str | None = None, + mobile_country_code: str | None = None, + mobile_network_code: str | None = None, + network: str | None = None, + organization: str | None = None, + static_ip_score: float | None = None, + user_type: UserType | None = None, + postal_code: str | None = None, + latitude: Decimal | None = None, + longitude: Decimal | None = None, + 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, + ) + + return _inner + + +@pytest.fixture +def user_payout_event_factory( + user_payout_event_manager: UserPayoutEventManager, +) -> Callable[..., UserPayoutEvent]: + + def _inner( + uuid: UUIDStr | None = None, + debit_account_uuid: UUIDStr | None = None, + account_reference_type: str | None = None, + account_reference_uuid: UUIDStr | None = None, + cashout_method_uuid: UUIDStr | None = None, + description: str | None = None, + created: AwareDatetimeISO | None = None, + amount: PositiveInt | None = None, + status: PayoutStatus | None = None, + ext_ref_id: str | None = None, + payout_type: PayoutType | None = None, + request_data: dict[str, Any] | None = None, + order_data: dict[str, Any] | CashMailOrderData | None = None, + ) -> UserPayoutEvent: + + debit_account_uuid = debit_account_uuid or uuid4().hex + cashout_method_uuid = cashout_method_uuid or uuid4().hex + # account_reference_type = account_reference_type or f"acct-ref-{uuid4().hex}" + # account_reference_uuid = account_reference_uuid or uuid4().hex + # cashout_method_uuid = cashout_method_uuid or uuid4().hex + amount = amount or randint(a=99, b=9_999) + status = status or rand_choice(list(PayoutStatus)) + + description = description or f"desc-{uuid4().hex[:12]}" + # ext_ref_id = ext_ref_id or f"ext-ref-{uuid4().hex[:8]}" + payout_type = payout_type or rand_choice(list(PayoutType)) + request_data = request_data or {} + # order_data = order_data or None + + return user_payout_event_manager.create( + uuid=uuid, + debit_account_uuid=debit_account_uuid, + account_reference_type=account_reference_type, + account_reference_uuid=account_reference_uuid, + cashout_method_uuid=cashout_method_uuid, + description=description, + created=created, + amount=amount, + status=status, + ext_ref_id=ext_ref_id, + payout_type=payout_type, + request_data=request_data, + order_data=order_data, + ) + + return _inner + + +@pytest.fixture +def iprecord_factory(iprecord_manager: IPRecordManager) -> Callable[..., IPRecord]: + + def _inner( + 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 iprecord_manager.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, + ) + + return _inner + + +# class AuditLogManager(PostgresManager): + + +@pytest.fixture +def auditlog_factory(audit_log_manager: AuditLogManager): + + 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: + + event_types = { + "offerwall-enter.blocked", + "offerwall-enter.rate-limited", + "offerwall-enter.url-modified", + } + + return audit_log_manager.create( + user_id=user_id, + level=level or rchoice(list(AuditLogLevel)), + event_type=event_type or rchoice(list(event_types)), + event_msg=event_msg, + event_value=event_value, + ) + + return _inner diff --git a/test_utils/models/upk/__init__.py b/test_utils/models/upk/__init__.py new file mode 100644 index 0000000..e69de29 --- /dev/null +++ b/test_utils/models/upk/__init__.py diff --git a/test_utils/models/upk/conftest.py b/test_utils/models/upk/conftest.py new file mode 100644 index 0000000..c8855da --- /dev/null +++ b/test_utils/models/upk/conftest.py @@ -0,0 +1,178 @@ +from __future__ import annotations + +import os +import time +from typing import TYPE_CHECKING +from uuid import UUID + +import pandas as pd +import pytest + +from generalresearch.pg_helper import PostgresConfig + +if TYPE_CHECKING: + from generalresearch.managers.thl.category import CategoryManager + + +def insert_data_from_csv( + thl_web_rw: PostgresConfig, + table_name: str, + fp: str | None = None, + disable_fk_checks: bool = False, + df: pd.DataFrame | None = None, +): + assert fp is not None or df is not None and not (fp is not None and df is not None) + if fp: + df = pd.read_csv(fp, dtype=str) + + assert isinstance(df, pd.DataFrame) + + df = df.where(pd.notnull(df), None) + cols = list(df.columns) + col_str = ", ".join(cols) + values_str = ", ".join(["%s"] * len(cols)) + if "id" in df.columns and len(df["id"].iloc[0]) == 36: + df["id"] = df["id"].map(lambda x: UUID(x).hex) + args = df.to_dict("tight")["data"] + + with thl_web_rw.make_connection() as conn: + with conn.cursor() as c: + if disable_fk_checks: + c.execute("SET CONSTRAINTS ALL DEFERRED") + c.executemany( + f"INSERT INTO {table_name} ({col_str}) VALUES ({values_str})", + params_seq=args, + ) + conn.commit() + + +@pytest.fixture(scope="session") +def category_data( + thl_web_rw: PostgresConfig, category_manager: CategoryManager +) -> None: + fp = os.path.join(os.path.dirname(__file__), "marketplace_category.csv.gz") + insert_data_from_csv( + thl_web_rw, + fp=fp, + table_name="marketplace_category", + disable_fk_checks=True, + ) + # Don't strictly need to do this, but probably we should + category_manager.populate_caches() + cats = category_manager.categories.values() + path_id = {c.path: c.id for c in cats} + data = [ + {"id": c.id, "parent_id": path_id[c.parent_path]} for c in cats if c.parent_path + ] + query = """ + UPDATE marketplace_category + SET parent_id = %(parent_id)s + WHERE id = %(id)s; + """ + with thl_web_rw.make_connection() as conn: + with conn.cursor() as c: + c.executemany(query=query, params_seq=data) + conn.commit() + + +@pytest.fixture(scope="session") +def property_data(thl_web_rw: PostgresConfig) -> None: + fp = os.path.join(os.path.dirname(__file__), "marketplace_property.csv.gz") + insert_data_from_csv(thl_web_rw, fp=fp, table_name="marketplace_property") + + +@pytest.fixture(scope="session") +def item_data(thl_web_rw: PostgresConfig) -> None: + fp = os.path.join(os.path.dirname(__file__), "marketplace_item.csv.gz") + insert_data_from_csv(thl_web_rw, fp=fp, table_name="marketplace_item") + + +@pytest.fixture(scope="session") +def propertycategoryassociation_data( + thl_web_rw: PostgresConfig, + category_data, + property_data, + category_manager: CategoryManager, +) -> None: + table_name = "marketplace_propertycategoryassociation" + fp = os.path.join(os.path.dirname(__file__), f"{table_name}.csv.gz") + # Need to lookup category pk from uuid + category_manager.populate_caches() + df = pd.read_csv(fp, dtype=str) + df["category_id"] = df["category_id"].map( + lambda x: category_manager.categories[x].id + ) + insert_data_from_csv(thl_web_rw, df=df, table_name=table_name) + + +@pytest.fixture(scope="session") +def propertycountry_data(thl_web_rw: PostgresConfig, property_data) -> None: + fp = os.path.join(os.path.dirname(__file__), "marketplace_propertycountry.csv.gz") + insert_data_from_csv(thl_web_rw, fp=fp, table_name="marketplace_propertycountry") + + +@pytest.fixture(scope="session") +def propertymarketplaceassociation_data( + thl_web_rw: PostgresConfig, property_data +) -> None: + table_name = "marketplace_propertymarketplaceassociation" + fp = os.path.join(os.path.dirname(__file__), f"{table_name}.csv.gz") + insert_data_from_csv(thl_web_rw, fp=fp, table_name=table_name) + + +@pytest.fixture(scope="session") +def propertyitemrange_data( + thl_web_rw: PostgresConfig, property_data, item_data +) -> None: + table_name = "marketplace_propertyitemrange" + fp = os.path.join(os.path.dirname(__file__), f"{table_name}.csv.gz") + insert_data_from_csv(thl_web_rw, fp=fp, table_name=table_name) + + +@pytest.fixture(scope="session") +def question_data(thl_web_rw: PostgresConfig) -> None: + table_name = "marketplace_question" + fp = os.path.join(os.path.dirname(__file__), f"{table_name}.csv.gz") + insert_data_from_csv( + thl_web_rw, fp=fp, table_name=table_name, disable_fk_checks=True + ) + + +@pytest.fixture(scope="session") +def clear_upk_tables(thl_web_rw: PostgresConfig): + tables = [ + "marketplace_propertyitemrange", + "marketplace_propertymarketplaceassociation", + "marketplace_propertycategoryassociation", + "marketplace_category", + "marketplace_item", + "marketplace_property", + "marketplace_propertycountry", + "marketplace_question", + ] + table_str = ", ".join(tables) + + with thl_web_rw.make_connection() as conn: + with conn.cursor() as c: + c.execute(f"TRUNCATE {table_str} RESTART IDENTITY CASCADE;") + conn.commit() + + +@pytest.fixture(scope="session") +def upk_data( + clear_upk_tables, + category_data, + property_data, + item_data, + propertycategoryassociation_data, + propertycountry_data, + propertymarketplaceassociation_data, + propertyitemrange_data, + question_data, +) -> None: + # Wait a second to make sure the HarmonizerCache refresh loop pulls these in + time.sleep(2) + + +def test_fixtures(upk_data): + pass diff --git a/test_utils/models/upk/marketplace_category.csv.gz b/test_utils/models/upk/marketplace_category.csv.gz Binary files differnew file mode 100644 index 0000000..0f8ec1c --- /dev/null +++ b/test_utils/models/upk/marketplace_category.csv.gz diff --git a/test_utils/models/upk/marketplace_item.csv.gz b/test_utils/models/upk/marketplace_item.csv.gz Binary files differnew file mode 100644 index 0000000..c12c5d8 --- /dev/null +++ b/test_utils/models/upk/marketplace_item.csv.gz diff --git a/test_utils/models/upk/marketplace_property.csv.gz b/test_utils/models/upk/marketplace_property.csv.gz Binary files differnew file mode 100644 index 0000000..a781d1d --- /dev/null +++ b/test_utils/models/upk/marketplace_property.csv.gz diff --git a/test_utils/models/upk/marketplace_propertycategoryassociation.csv.gz b/test_utils/models/upk/marketplace_propertycategoryassociation.csv.gz Binary files differnew file mode 100644 index 0000000..5b4ea19 --- /dev/null +++ b/test_utils/models/upk/marketplace_propertycategoryassociation.csv.gz diff --git a/test_utils/models/upk/marketplace_propertycountry.csv.gz b/test_utils/models/upk/marketplace_propertycountry.csv.gz Binary files differnew file mode 100644 index 0000000..5d2a637 --- /dev/null +++ b/test_utils/models/upk/marketplace_propertycountry.csv.gz diff --git a/test_utils/models/upk/marketplace_propertyitemrange.csv.gz b/test_utils/models/upk/marketplace_propertyitemrange.csv.gz Binary files differnew file mode 100644 index 0000000..84f4f0e --- /dev/null +++ b/test_utils/models/upk/marketplace_propertyitemrange.csv.gz diff --git a/test_utils/models/upk/marketplace_propertymarketplaceassociation.csv.gz b/test_utils/models/upk/marketplace_propertymarketplaceassociation.csv.gz Binary files differnew file mode 100644 index 0000000..6b9fd1c --- /dev/null +++ b/test_utils/models/upk/marketplace_propertymarketplaceassociation.csv.gz diff --git a/test_utils/models/upk/marketplace_question.csv.gz b/test_utils/models/upk/marketplace_question.csv.gz Binary files differnew file mode 100644 index 0000000..bcfc3ad --- /dev/null +++ b/test_utils/models/upk/marketplace_question.csv.gz |
