From e2c5de703be45746bacaea4136f24440ff5a291c Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Mon, 24 Aug 2026 12:35:35 -0700 Subject: Ruff std replacements --- tests/incite/test_interval_idx.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) (limited to 'tests/incite/test_interval_idx.py') diff --git a/tests/incite/test_interval_idx.py b/tests/incite/test_interval_idx.py index ea2bced..3034c21 100644 --- a/tests/incite/test_interval_idx.py +++ b/tests/incite/test_interval_idx.py @@ -1,5 +1,6 @@ +from datetime import datetime + import pandas as pd -from datetime import datetime, timezone, timedelta class TestIntervalIndex: -- cgit v1.2.3 From aeeb7fef2594ccd34fbe96a77f6c5b392299fed7 Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Thu, 27 Aug 2026 17:46:11 -0700 Subject: Ruff afternoon --- generalresearch/managers/thl/product.py | 1 - generalresearch/models/network/nmap/parser.py | 8 +- generalresearch/models/precision/question.py | 4 +- generalresearch/models/precision/survey.py | 17 +- generalresearch/models/prodege/survey.py | 3 +- generalresearch/models/spectrum/question.py | 9 +- generalresearch/models/spectrum/survey.py | 13 +- generalresearch/models/spectrum/task_collection.py | 2 +- generalresearch/models/thl/contest/contest.py | 10 +- generalresearch/models/thl/contest/leaderboard.py | 13 +- generalresearch/models/thl/contest/milestone.py | 20 +- generalresearch/models/thl/contest/raffle.py | 46 ++-- generalresearch/models/thl/finance.py | 12 +- generalresearch/models/thl/product.py | 36 +-- .../models/thl/profiling/other_option.py | 4 +- .../models/thl/profiling/upk_question.py | 13 +- .../models/thl/profiling/user_question_answer.py | 31 +-- pyproject.toml | 5 +- .../collections/test_df_collection_item_thl_web.py | 273 ++++++++++----------- .../mergers/foundations/test_enriched_session.py | 52 ++-- .../foundations/test_enriched_task_adjust.py | 38 ++- .../mergers/foundations/test_enriched_wall.py | 73 +++--- tests/incite/mergers/test_merge_collection.py | 53 +++- tests/incite/mergers/test_merge_collection_item.py | 25 +- tests/incite/mergers/test_pop_ledger.py | 109 ++++---- tests/incite/mergers/test_ym_survey_merge.py | 55 +++-- tests/incite/schemas/test_admin_responses.py | 36 +-- tests/incite/schemas/test_thl_web.py | 8 +- tests/incite/test_collection_base.py | 62 ++--- tests/incite/test_collection_base_item.py | 74 +++--- tests/incite/test_interval_idx.py | 6 +- tests/managers/gr/test_business.py | 60 +++-- tests/managers/gr/test_team.py | 58 +++-- tests/managers/network/test_label.py | 4 +- .../managers/thl/test_contest/test_leaderboard.py | 5 +- tests/managers/thl/test_contest/test_milestone.py | 5 +- tests/managers/thl/test_contest/test_raffle.py | 12 +- tests/managers/thl/test_ledger/test_lm_accounts.py | 2 +- tests/managers/thl/test_ledger/test_lm_tx.py | 4 - .../thl/test_ledger/test_thl_lm_accounts.py | 11 +- .../thl/test_ledger/test_thl_lm_bp_payout.py | 54 ++-- tests/managers/thl/test_ledger/test_thl_lm_tx.py | 26 +- tests/managers/thl/test_payout.py | 2 - tests/managers/thl/test_survey.py | 6 +- tests/managers/thl/test_user_manager/test_base.py | 2 +- tests/models/network/test_nmap.py | 2 +- tests/models/spectrum/test_survey.py | 44 ++++ tests/models/thl/test_product.py | 4 +- 48 files changed, 815 insertions(+), 597 deletions(-) (limited to 'tests/incite/test_interval_idx.py') diff --git a/generalresearch/managers/thl/product.py b/generalresearch/managers/thl/product.py index d924e17..54fa7c8 100644 --- a/generalresearch/managers/thl/product.py +++ b/generalresearch/managers/thl/product.py @@ -33,7 +33,6 @@ if TYPE_CHECKING: ProfilingConfig, SessionConfig, SourcesConfig, - SupplyConfigs, UserCreateConfig, UserHealthConfig, UserWalletConfig, diff --git a/generalresearch/models/network/nmap/parser.py b/generalresearch/models/network/nmap/parser.py index ecaf2d1..866b4bd 100644 --- a/generalresearch/models/network/nmap/parser.py +++ b/generalresearch/models/network/nmap/parser.py @@ -48,7 +48,7 @@ class NmapXmlParser: try: root = ET.fromstring(nmap_data) - except Exception as e: + except ET.ParseError as e: emsg = f"Wrong XML structure: cannot parse data: {e}" raise NmapParserException(emsg) @@ -103,7 +103,7 @@ class NmapXmlParser: @classmethod def _parse_scaninfo(cls, scaninfo_el: ET.Element) -> NmapScanInfo: - data = dict() + data = {} data["type"] = NmapScanType(scaninfo_el.attrib["type"]) data["protocol"] = IPProtocol(scaninfo_el.attrib["protocol"]) data["num_services"] = scaninfo_el.attrib["numservices"] @@ -132,7 +132,7 @@ class NmapXmlParser: @classmethod def _parse_nmaprun(cls, nmaprun_el: ET.Element) -> dict: - nmap_data = dict() + nmap_data = {} nmaprun = dict(nmaprun_el.attrib) nmap_data["command_line"] = nmaprun["args"] nmap_data["started_at"] = datetime.fromtimestamp( @@ -148,7 +148,7 @@ class NmapXmlParser: Receives a XML tag representing a scanned host with its services. """ - data = dict() + data = {} # status_el = host_el.find("status") diff --git a/generalresearch/models/precision/question.py b/generalresearch/models/precision/question.py index a2189d5..cc90aa9 100644 --- a/generalresearch/models/precision/question.py +++ b/generalresearch/models/precision/question.py @@ -6,7 +6,7 @@ import logging from enum import StrEnum from typing import TYPE_CHECKING, Any, Literal -from pydantic import BaseModel, Field, field_validator, model_validator +from pydantic import BaseModel, Field, ValidationError, field_validator, model_validator from generalresearch.models import Source, string_utils from generalresearch.models.precision import PrecisionQuestionID @@ -112,7 +112,7 @@ class PrecisionQuestion(MarketplaceQuestion): """ try: return cls._from_api(d) - except Exception as e: + except ValidationError as e: logger.warning(f"Unable to parse question: {d}. {e}") return None diff --git a/generalresearch/models/precision/survey.py b/generalresearch/models/precision/survey.py index f515552..b27b8c4 100644 --- a/generalresearch/models/precision/survey.py +++ b/generalresearch/models/precision/survey.py @@ -99,11 +99,11 @@ class PrecisionQuota(BaseModel): self, criteria_evaluation: dict[str, bool | None] ) -> tuple[bool | None, list[str]]: # Passes back "matches" (T/F/none) and a list of unknown criterion hashes - unknowns = list() + unknowns = [] for c in self.condition_hashes: eval_value = criteria_evaluation.get(c) if eval_value is False: - return False, list() + return False, [] if eval_value is None: unknowns.append(c) if unknowns: @@ -245,11 +245,10 @@ class PrecisionSurvey(MarketplaceTask): # Fancy repr that abbreviates exclude_pids and excluded_surveys repr_args = list(self.__repr_args__()) for n, (k, v) in enumerate(repr_args): - if k in {"excluded_surveys"}: - if v and len(v) > 6: - v = sorted(v) - v = v[:3] + ["…"] + v[-3:] - repr_args[n] = (k, v) + if k in {"excluded_surveys"} and v and len(v) > 6: + v = sorted(v) + v = v[:3] + ["…"] + v[-3:] + repr_args[n] = (k, v) join_str = ", " repr_str = join_str.join( repr(v) if a is None else f"{a}={v!r}" for a, v in repr_args @@ -369,6 +368,4 @@ class PrecisionSurvey(MarketplaceTask): return False if self.group_id in att_group_ids: return False - if self.excluded_surveys & att_survey_ids: - return False - return True + return not self.excluded_surveys & att_survey_ids diff --git a/generalresearch/models/prodege/survey.py b/generalresearch/models/prodege/survey.py index 3f4c88f..7ab6df6 100644 --- a/generalresearch/models/prodege/survey.py +++ b/generalresearch/models/prodege/survey.py @@ -13,6 +13,7 @@ from pydantic import ( BaseModel, ConfigDict, Field, + ValidationError, computed_field, field_validator, model_validator, @@ -513,7 +514,7 @@ class ProdegeSurvey(MarketplaceTask): def from_api(cls, d: dict[str, Any]) -> ProdegeSurvey | None: try: return cls._from_api(d) - except Exception as e: + except ValidationError as e: logger.warning(f"Unable to parse survey: {d}. {e}") return None diff --git a/generalresearch/models/spectrum/question.py b/generalresearch/models/spectrum/question.py index c8eea4a..7add692 100644 --- a/generalresearch/models/spectrum/question.py +++ b/generalresearch/models/spectrum/question.py @@ -13,6 +13,7 @@ from pydantic import ( BaseModel, Field, PositiveInt, + ValidationError, field_validator, model_validator, ) @@ -132,7 +133,7 @@ class SpectrumQuestionType(StrEnum): @classmethod def from_api(cls, a: int): api_type_map = cls.get_api_map() - return api_type_map[a] if a in api_type_map else None + return api_type_map.get(a, None) class SpectrumQuestionClass(IntEnum): @@ -260,7 +261,7 @@ class SpectrumQuestion(MarketplaceQuestion): return None try: return cls._from_api(d, country_iso, language_iso) - except Exception as e: + except ValidationError as e: logger.warning(f"Unable to parse question: {d}. {e}") return None @@ -280,7 +281,9 @@ class SpectrumQuestion(MarketplaceQuestion): ] created = ( - datetime.utcfromtimestamp(d["crtd_on"] / 1000).replace(tzinfo=UTC) + datetime.fromtimestamp(timestamp=d["crtd_on"] / 1000, tz=UTC).replace( + tzinfo=UTC + ) if d.get("crtd_on") else None ) diff --git a/generalresearch/models/spectrum/survey.py b/generalresearch/models/spectrum/survey.py index f9a5e27..424d206 100644 --- a/generalresearch/models/spectrum/survey.py +++ b/generalresearch/models/spectrum/survey.py @@ -75,8 +75,7 @@ class SpectrumCondition(MarketplaceCondition): rs["from"] = round(rs["from"] / 12) rs["to"] = round(rs["to"] / 12) d["values"] = [ - f"{rs["from"] or "inf"}-{rs["to"] or "inf"}" - for rs in d["range_sets"] + f"{rs["from"] or "inf"}-{rs["to"] or "inf"}" for rs in d["range_sets"] ] d["value_type"] = ConditionValueType.RANGE return cls.model_validate(d) @@ -103,7 +102,7 @@ class SpectrumQuota(BaseModel): # There is no explicit status. The quota is closed if the count is 0 def __hash__(self) -> int: - return hash(tuple((tuple(self.condition_hashes), self.remaining_count))) + return hash((tuple(self.condition_hashes), self.remaining_count)) @property def is_open(self) -> bool: @@ -113,7 +112,7 @@ class SpectrumQuota(BaseModel): return self.remaining_count >= min_open_spots @classmethod - def from_api(cls, d: dict) -> Self: + def from_api(cls, d: dict[str, Any]) -> Self: d["remaining_count"] = d["quantities"]["currently_open"] return cls.model_validate(d) @@ -323,7 +322,7 @@ class SpectrumSurvey(MarketplaceTask): def from_api(cls, d: dict[str, Any]) -> SpectrumSurvey | None: try: return cls._from_api(d) - except Exception as e: + except (AssertionError, ValueError) as e: logger.warning(f"Unable to parse survey: {d}. {e}") return None @@ -336,7 +335,7 @@ class SpectrumSurvey(MarketplaceTask): else TaskCalculationType.COMPLETES ) - d["conditions"] = dict() + d["conditions"] = {} # If we haven't hit the "detail" endpoint, we won't get this d.setdefault("qualifications", []) @@ -454,7 +453,7 @@ class SpectrumSurvey(MarketplaceTask): quota_eval = { quota: quota.matches_soft(criteria_evaluation) for quota in self.quotas } - evals = set(g[0] for g in quota_eval.values()) + evals = {g[0] for g in quota_eval.values()} if any(m[0] is True and not q.is_open for q, m in quota_eval.items()): # matched a full quota return False, set() diff --git a/generalresearch/models/spectrum/task_collection.py b/generalresearch/models/spectrum/task_collection.py index 6715378..8ca5a93 100644 --- a/generalresearch/models/spectrum/task_collection.py +++ b/generalresearch/models/spectrum/task_collection.py @@ -91,7 +91,7 @@ class SpectrumTaskCollection(TaskCollection): "survey_id", ] rows = [] - d = dict() + d = {} for k in fields: d[k] = getattr(s, k) if hasattr(s, k) else None d["used_question_ids"] = list(s.used_question_ids) diff --git a/generalresearch/models/thl/contest/contest.py b/generalresearch/models/thl/contest/contest.py index 2a8853d..bd0fc04 100644 --- a/generalresearch/models/thl/contest/contest.py +++ b/generalresearch/models/thl/contest/contest.py @@ -136,10 +136,12 @@ class Contest(ContestBase): # return False def should_end(self) -> tuple[bool, ContestEndReason | None]: - if self.status == ContestStatus.ACTIVE: - if self.end_condition.ends_at: - if datetime.now(tz=UTC) >= self.end_condition.ends_at: - return True, ContestEndReason.ENDS_AT + if ( + self.status == ContestStatus.ACTIVE + and self.end_condition.ends_at + and datetime.now(tz=UTC) >= self.end_condition.ends_at + ): + return True, ContestEndReason.ENDS_AT return False, None diff --git a/generalresearch/models/thl/contest/leaderboard.py b/generalresearch/models/thl/contest/leaderboard.py index 696cdea..e923383 100644 --- a/generalresearch/models/thl/contest/leaderboard.py +++ b/generalresearch/models/thl/contest/leaderboard.py @@ -151,7 +151,8 @@ class LeaderboardContest(LeaderboardContestCreate, Contest): len(self.country_isos) == 1 ), "Can only set 1 country_iso in a leaderboard contest" assert ( - list(self.country_isos)[0] == self.leaderboard_key_parts["country_iso"] + next(iter(self.country_isos)) + == self.leaderboard_key_parts["country_iso"] ), "leaderboard_key country_iso must match the country_isos" else: self.country_isos = {self.leaderboard_key_parts["country_iso"]} @@ -192,10 +193,12 @@ class LeaderboardContest(LeaderboardContestCreate, Contest): return lbm def should_end(self) -> tuple[bool, ContestEndReason | None]: - if self.status == ContestStatus.ACTIVE: - if self.end_condition.ends_at: - if datetime.now(tz=UTC) >= self.end_condition.ends_at: - return True, ContestEndReason.ENDS_AT + if ( + self.status == ContestStatus.ACTIVE + and self.end_condition.ends_at + and datetime.now(tz=UTC) >= self.end_condition.ends_at + ): + return True, ContestEndReason.ENDS_AT return False, None diff --git a/generalresearch/models/thl/contest/milestone.py b/generalresearch/models/thl/contest/milestone.py index 8d96fcb..5fc27fa 100644 --- a/generalresearch/models/thl/contest/milestone.py +++ b/generalresearch/models/thl/contest/milestone.py @@ -132,10 +132,12 @@ class MilestoneContest(MilestoneContestCreate, Contest): if res: return res, msg - if self.status == ContestStatus.ACTIVE: - if self.end_condition.max_winners: - if self.win_count >= self.end_condition.max_winners: - return True, ContestEndReason.MAX_WINNERS + if ( + self.status == ContestStatus.ACTIVE + and self.end_condition.max_winners + and self.win_count >= self.end_condition.max_winners + ): + return True, ContestEndReason.MAX_WINNERS return False, None @@ -189,16 +191,10 @@ class MilestoneUserView(MilestoneContest, ContestUserView): ) def should_award(self): - if self.status == ContestStatus.ACTIVE: - if self.should_have_awarded(): - return True - return False + return bool(self.status == ContestStatus.ACTIVE and self.should_have_awarded()) def should_have_awarded(self): - if self.target_amount: - if self.user_amount >= self.target_amount: - return True - return False + return bool(self.target_amount and self.user_amount >= self.target_amount) def is_user_eligible(self, country_iso: str) -> tuple[bool, str]: passes, msg = super().is_user_eligible(country_iso=country_iso) diff --git a/generalresearch/models/thl/contest/raffle.py b/generalresearch/models/thl/contest/raffle.py index 08243f4..16a0a47 100644 --- a/generalresearch/models/thl/contest/raffle.py +++ b/generalresearch/models/thl/contest/raffle.py @@ -127,7 +127,7 @@ class RaffleContest(RaffleContestCreate, Contest): # If there is more than 1 prize, the winning entry is subtracted # from the user's entry count user_amount = defaultdict(int) - user_id_user = dict() + user_id_user = {} for entry in self.entries: user_amount[entry.user.user_id] += entry.amount user_id_user[entry.user.user_id] = entry.user @@ -149,10 +149,12 @@ class RaffleContest(RaffleContestCreate, Contest): res, msg = super().should_end() if res: return res, msg - if self.status == ContestStatus.ACTIVE: - if self.end_condition.target_entry_amount: - if self.current_amount >= self.end_condition.target_entry_amount: - return True, ContestEndReason.TARGET_ENTRY_AMOUNT + if ( + self.status == ContestStatus.ACTIVE + and self.end_condition.target_entry_amount + and self.current_amount >= self.end_condition.target_entry_amount + ): + return True, ContestEndReason.TARGET_ENTRY_AMOUNT return False, None @staticmethod @@ -278,17 +280,19 @@ class RaffleUserView(RaffleContest, ContestUserView): return probs def is_entry_eligible(self, entry: ContestEntry) -> tuple[bool, str]: - if self.entry_rule.max_entry_amount_per_user: - if ( - self.user_amount + entry.amount - ) > self.entry_rule.max_entry_amount_per_user: - return False, "Entry would exceed max amount per user." - - if self.entry_rule.max_daily_entries_per_user: - if ( - self.user_amount_today + entry.amount - ) > self.entry_rule.max_daily_entries_per_user: - return False, "Entry would exceed max amount per user per day." + if ( + self.entry_rule.max_entry_amount_per_user + and (self.user_amount + entry.amount) + > self.entry_rule.max_entry_amount_per_user + ): + return False, "Entry would exceed max amount per user." + + if ( + self.entry_rule.max_daily_entries_per_user + and (self.user_amount_today + entry.amount) + > self.entry_rule.max_daily_entries_per_user + ): + return False, "Entry would exceed max amount per user per day." return True, "" def is_user_eligible(self, country_iso: str) -> tuple[bool, str]: @@ -296,16 +300,18 @@ class RaffleUserView(RaffleContest, ContestUserView): if not passes: return False, msg - if self.entry_rule.max_entry_amount_per_user: + if self.entry_rule.max_entry_amount_per_user: # noqa: SIM102 # Greater or equal b/c we're asking if the user is eligible to # enter MORE, now! If it equals, nothing is wrong, just that they # are not eligible anymore. if self.user_amount >= self.entry_rule.max_entry_amount_per_user: return False, "Reached max amount per user." - if self.entry_rule.max_daily_entries_per_user: - if self.user_amount_today >= self.entry_rule.max_daily_entries_per_user: - return False, "Reached max amount today." + if ( + self.entry_rule.max_daily_entries_per_user + and self.user_amount_today >= self.entry_rule.max_daily_entries_per_user + ): + return False, "Reached max amount today." # This would indicate something is wrong, as something else should have done this e, _ = self.should_end() diff --git a/generalresearch/models/thl/finance.py b/generalresearch/models/thl/finance.py index 0856825..8c94390 100644 --- a/generalresearch/models/thl/finance.py +++ b/generalresearch/models/thl/finance.py @@ -124,10 +124,10 @@ class POPFinancial(BaseModel): Direction, ) - assert all([a.account_type == AccountType.BP_WALLET for a in accounts]) - assert all([a.normal_balance == Direction.CREDIT for a in accounts]) + assert all(a.account_type == AccountType.BP_WALLET for a in accounts) + assert all(a.normal_balance == Direction.CREDIT for a in accounts) if not is_debug(): - assert all([a.currency == "USD" for a in accounts]) + assert all(a.currency == "USD" for a in accounts) if input_data.empty: return [] @@ -850,11 +850,11 @@ class BusinessBalances(BaseModel): # Validate the input accounts assert len(accounts) > 0, "Must provide accounts" - assert all([a.account_type == AccountType.BP_WALLET for a in accounts]) - assert all([a.normal_balance == Direction.CREDIT for a in accounts]) + assert all(a.account_type == AccountType.BP_WALLET for a in accounts) + assert all(a.normal_balance == Direction.CREDIT for a in accounts) if not is_debug(): - assert all([a.currency == "USD" for a in accounts]) + assert all(a.currency == "USD" for a in accounts) # Validate the input dataframe assert input_data.index.name == "account_id" diff --git a/generalresearch/models/thl/product.py b/generalresearch/models/thl/product.py index 76a8e83..a7ecd55 100644 --- a/generalresearch/models/thl/product.py +++ b/generalresearch/models/thl/product.py @@ -430,8 +430,8 @@ class UserWalletConfig(BaseModel): @field_serializer("supported_payout_types", when_used="json") def serialize_supported_payout_types_in_order( self, supported_payout_types: set[PayoutType] - ) -> set[PayoutType]: - return set(sorted(supported_payout_types)) + ) -> list[PayoutType]: + return sorted(supported_payout_types) @field_validator("min_cashout", mode="after") @classmethod @@ -552,14 +552,14 @@ class PayoutTransformation(BaseModel): min_payout = Decimal(0) pct = Decimal(pct) - payout = Decimal(payout) + _payout = Decimal(payout) min_payout = Decimal(min_payout) max_payout = Decimal(max_payout) if max_payout else None - payout: Decimal = payout * pct - payout: Decimal = max([payout, min_payout]) - payout: Decimal = min([payout, max_payout]) if max_payout else payout - return payout + _payout: Decimal = _payout * pct + _payout: Decimal = max([_payout, min_payout]) + _payout: Decimal = min([_payout, max_payout]) if max_payout else payout + return _payout def payout_transformation_amt( self, payout: Decimal, user_wallet_balance: Decimal | None = None @@ -569,22 +569,22 @@ class PayoutTransformation(BaseModel): # (display, adjustment) so ignore the 7-cent rounding. if user_wallet_balance is None: return self.payout_transformation_percent(payout=payout, pct=Decimal(".95")) - payout = Decimal(payout) + _payout = Decimal(payout) - payout: Decimal = payout * Decimal("0.95") - new_balance = payout + user_wallet_balance + _payout: Decimal = _payout * Decimal("0.95") + new_balance = _payout + user_wallet_balance # If the new_balance is <0, we aren't paying anything, so use the # full amount if new_balance < 0: - return payout + return _payout amt = (5 * math.floor((int(new_balance * 100) - 2) / 5)) + 2 rounded_new_balance = Decimal(amt / 100).quantize(Decimal("0.00")) - payout = rounded_new_balance - user_wallet_balance - if payout < Decimal(0): + _payout = rounded_new_balance - user_wallet_balance + if _payout < Decimal(0): return Decimal(0) - return payout + return _payout class SourceConfig(BaseModel): @@ -731,8 +731,8 @@ class SupplyConfig(BaseModel): Use global config. """ d = self.global_scoped_policies_dict.copy() - d.update(self.team_scoped_policies_dict.get(team_id, dict())) - d.update(self.product_scoped_policies_dict.get(product_id, dict())) + d.update(self.team_scoped_policies_dict.get(team_id, {})) + d.update(self.product_scoped_policies_dict.get(product_id, {})) return d def get_config_for_product(self, product: Product) -> MergedSupplyConfig: @@ -751,7 +751,7 @@ class SupplyConfig(BaseModel): supply_policy=policy_dict[source], source_config=sources_dict[source], ) - for source in policy_dict.keys() + for source in policy_dict ] ) @@ -1000,10 +1000,12 @@ class Product(BaseModel, validate_assignment=True): @property def business_uuid(self) -> UUIDStr: + assert self.business_id return self.business_id @property def team_uuid(self) -> UUIDStr: + assert self.team_id return self.team_id @property diff --git a/generalresearch/models/thl/profiling/other_option.py b/generalresearch/models/thl/profiling/other_option.py index 6d789e5..2f3cac9 100644 --- a/generalresearch/models/thl/profiling/other_option.py +++ b/generalresearch/models/thl/profiling/other_option.py @@ -51,6 +51,4 @@ def option_is_catch_all(c: UpkQuestionChoice) -> bool: return True if c.text.lower() in texts_exact: return True - if any(t in c.text.lower() for t in texts_in): - return True - return False + return bool(any(t in c.text.lower() for t in texts_in)) diff --git a/generalresearch/models/thl/profiling/upk_question.py b/generalresearch/models/thl/profiling/upk_question.py index 77bba6f..3bb0733 100644 --- a/generalresearch/models/thl/profiling/upk_question.py +++ b/generalresearch/models/thl/profiling/upk_question.py @@ -475,10 +475,9 @@ class UpkQuestion(BaseModel): # Almost nothing has >1k options, besides location stuff (cities, # etc.) which should get harmonized. When presenting them, we'll # filter down options to at most 50. - if self.choices and (len(self.choices) <= 1 or len(self.choices) > 1000): - return False - - return True + return not ( + self.choices and (len(self.choices) <= 1 or len(self.choices) > 1000) + ) @property def md5sum(self): @@ -534,7 +533,7 @@ class UpkQuestion(BaseModel): ), "Multiple of the same answer submitted" if self.type == UpkQuestionType.MULTIPLE_CHOICE: assert len(answer) >= 1, "MC question with no selected answers" - choice_codes = set(x.id for x in self.choices) + choice_codes = {x.id for x in self.choices} if self.selector == UpkQuestionSelectorMC.SINGLE_ANSWER: assert ( len(answer) == 1 @@ -563,9 +562,7 @@ class UpkQuestion(BaseModel): assert len(answer) == 1, "Only one answer allowed" answer = answer[0] assert len(answer) > 0, "Must provide answer" - max_length = ( - self.configuration.max_length if self.configuration else 0 or 100000 - ) + max_length = self.configuration.max_length if self.configuration else 100000 assert len(answer) <= max_length, "Answer longer than allowed" if self.validation and self.validation.patterns: for pattern in self.validation.patterns: diff --git a/generalresearch/models/thl/profiling/user_question_answer.py b/generalresearch/models/thl/profiling/user_question_answer.py index 2db07b7..378345e 100644 --- a/generalresearch/models/thl/profiling/user_question_answer.py +++ b/generalresearch/models/thl/profiling/user_question_answer.py @@ -3,7 +3,7 @@ from __future__ import annotations import json from collections.abc import Iterator from datetime import UTC, datetime, timedelta -from typing import Any, Literal, Self +from typing import Any, Literal from pydantic import ( BaseModel, @@ -14,7 +14,6 @@ from pydantic import ( model_validator, ) -from generalresearch.grpc import timestamp_to_datetime from generalresearch.models import MAX_INT32, Source from generalresearch.models.custom_types import AwareDatetimeISO, UUIDStr from generalresearch.models.thl.locales import CountryISO, LanguageISO @@ -39,19 +38,23 @@ class UserQuestionAnswer(BaseModel): calc_answers: dict[str, tuple[str, ...]] | None = Field(default=None) @field_validator("calc_answers") - def sorted_calc_answers(cls, calc_answers) -> dict[str, tuple[str, ...]] | None: + def sorted_calc_answers( + cls, calc_answers: dict[str, tuple[str, ...]] | None + ) -> dict[str, tuple[str, ...]] | None: if calc_answers is None: return None return {k: tuple(sorted(v)) for k, v in calc_answers.items()} @field_validator("calc_answers") - def validate_keys(cls, calc_answers) -> dict[str, tuple[str, ...]] | None: + def validate_keys( + cls, calc_answers: dict[str, tuple[str, ...]] | None + ) -> dict[str, tuple[str, ...]] | None: if calc_answers is None: return None assert all( - ":" in k for k in calc_answers.keys() + ":" in k for k in calc_answers ), "calc_answers expects the keys to be in format source:question_code" return calc_answers @@ -66,6 +69,7 @@ class UserQuestionAnswer(BaseModel): return d def get_mrpqs(self) -> Iterator[MarketplaceResearchProfileQuestion]: + assert self.calc_answers for k, v in self.calc_answers.items(): source, question_code = k.split(":", 1) yield MarketplaceResearchProfileQuestion( @@ -105,21 +109,6 @@ class UserQuestionAnswer(BaseModel): def is_stale(self) -> bool: return self.timestamp < datetime.now(tz=UTC) - timedelta(days=30) - @classmethod - def from_grpc(cls, msg, default_timestamp: datetime) -> Self: - """ - Handles correctly issues with grpc timestamps - :param msg: "thl.protos.generalresearch_pb2.ProfilingQuestionAnswer" - """ - assert default_timestamp.tzinfo is not None, "must use tz-aware timestamps" - timestamp = timestamp_to_datetime(msg.timestamp) - timestamp = default_timestamp if timestamp < datetime(2000, 1, 1) else timestamp - return cls( - question_id=msg.question_id, - answer=tuple(msg.answer), - timestamp=timestamp, - ) - # We can't set a redis list to [] vs None. We'll push this dummy answer into # the cache to signify the user has no answered questions. It'll get removed @@ -131,7 +120,7 @@ DUMMY_UQA = UserQuestionAnswer( country_iso="xx", language_iso="xxx", property_code="dummy", - calc_answers=dict(), + calc_answers={}, ) diff --git a/pyproject.toml b/pyproject.toml index dbdf3b9..03a1a1f 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -56,4 +56,7 @@ testpaths = ["tests"] addopts = "-v --tb=short" [tool.ruff] -target-version = "py314" \ No newline at end of file +target-version = "py314" +exclude = [ + "generalresearch/thl_django", +] \ No newline at end of file diff --git a/tests/incite/collections/test_df_collection_item_thl_web.py b/tests/incite/collections/test_df_collection_item_thl_web.py index 8038d3b..edf90f7 100644 --- a/tests/incite/collections/test_df_collection_item_thl_web.py +++ b/tests/incite/collections/test_df_collection_item_thl_web.py @@ -5,12 +5,12 @@ from datetime import UTC, datetime, timedelta from itertools import product as iter_product from os.path import join as pjoin from pathlib import Path, PurePath -from typing import TYPE_CHECKING from uuid import uuid4 import dask.dataframe as dd import pandas as pd import pytest +from dask.distributed import Client as DaskClient from distributed import Client, Scheduler, Worker # noinspection PyUnresolvedReferences @@ -21,20 +21,19 @@ from faker import Faker from pandera.pandas import DataFrameSchema from pydantic import FilePath -from generalresearch.incite.base import CollectionItemBase +from generalresearch.incite.base import CollectionItemBase, GRLDatasets from generalresearch.incite.collections import ( + DFCollection, DFCollectionItem, DFCollectionType, ) from generalresearch.incite.schemas import ARCHIVE_AFTER +from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager from generalresearch.models.thl.product import Product from generalresearch.models.thl.user import User from generalresearch.pg_helper import PostgresConfig from generalresearch.sql_helper import PostgresDsn -if TYPE_CHECKING: - from generalresearch.incite.base import GRLDatasets - fake = Faker() df_collections = [ @@ -72,7 +71,12 @@ class TestDFCollectionItemBase: ) class TestDFCollectionItemProperties: - def test_filename(self, df_collection_data_type, df_collection, offset: str): + def test_filename( + self, + df_collection_data_type: DFCollectionType, + df_collection: DFCollection, + offset: str, + ): for i in df_collection.items: assert isinstance(i.filename, str) @@ -89,37 +93,59 @@ class TestDFCollectionItemProperties: ) class TestDFCollectionItemPropertiesBase: - def test_name(self, df_collection_data_type, offset: str, df_collection): + def test_name( + self, + df_collection: DFCollection, + ): for i in df_collection.items: assert isinstance(i.name, str) - def test_finish(self, df_collection_data_type, offset: str, df_collection): + def test_finish( + self, + df_collection: DFCollection, + ): for i in df_collection.items: assert isinstance(i.finish, datetime) - def test_interval(self, df_collection_data_type, offset: str, df_collection): + def test_interval( + self, + df_collection: DFCollection, + ): for i in df_collection.items: assert isinstance(i.interval, pd.Interval) def test_partial_filename( - self, df_collection_data_type, offset: str, df_collection + self, + df_collection: DFCollection, ): for i in df_collection.items: assert isinstance(i.partial_filename, str) - def test_empty_filename(self, df_collection_data_type, offset: str, df_collection): + def test_empty_filename( + self, + df_collection: DFCollection, + ): for i in df_collection.items: assert isinstance(i.empty_filename, str) - def test_path(self, df_collection_data_type, offset: str, df_collection): + def test_path( + self, + df_collection: DFCollection, + ): for i in df_collection.items: assert isinstance(i.path, FilePath) - def test_partial_path(self, df_collection_data_type, offset: str, df_collection): + def test_partial_path( + self, + df_collection: DFCollection, + ): for i in df_collection.items: assert isinstance(i.partial_path, FilePath) - def test_empty_path(self, df_collection_data_type, offset: str, df_collection): + def test_empty_path( + self, + df_collection: DFCollection, + ): for i in df_collection.items: assert isinstance(i.empty_path, FilePath) @@ -138,11 +164,8 @@ class TestDFCollectionItemMethod: def test_has_mysql( self, - df_collection, + df_collection: DFCollection, thl_web_rr: PostgresConfig, - offset: str, - duration: timedelta, - df_collection_data_type, delete_df_collection: Callable[..., None], ): delete_df_collection(coll=df_collection) @@ -168,12 +191,6 @@ class TestDFCollectionItemMethod: @pytest.mark.skip def test_update_partial_archive( self, - df_collection, - offset: str, - duration: timedelta, - thl_web_rw: PostgresConfig, - df_collection_data_type, - delete_df_collection: Callable[..., None], ): # for i in collection.items: # assert i.update_partial_archive() @@ -183,28 +200,12 @@ class TestDFCollectionItemMethod: @pytest.mark.skip def test_create_partial_archive( self, - df_collection, - offset: str, - duration: str, - create_main_accounts: Callable[..., None], - thl_web_rw: PostgresConfig, - thl_lm, - df_collection_data_type, - user_factory: Callable[..., User], - product: product: Product, - client_no_amm, - incite_item_factory, - delete_df_collection: Callable[..., None], - mnt_filepath: GRLDatasets, ): assert 1 + 1 == 2 def test_dict( self, - df_collection_data_type, - offset: str, - duration: timedelta, - df_collection, + df_collection: DFCollection, delete_df_collection: Callable[..., None], ): delete_df_collection(coll=df_collection) @@ -225,15 +226,15 @@ class TestDFCollectionItemMethod: def test_from_mysql( self, - df_collection_data_type, - df_collection, + df_collection_data_type: DFCollectionType, + df_collection: DFCollection, offset: str, duration: timedelta, create_main_accounts: Callable[..., None], thl_web_rw: PostgresConfig, user_factory: Callable[..., User], - product: product: Product, - incite_item_factory, + product: Product, + incite_item_factory: Callable[..., None], delete_df_collection: Callable[..., None], ): @@ -253,12 +254,14 @@ class TestDFCollectionItemMethod: if df_collection.data_type == DFCollectionType.LEDGER: assert df is None else: + assert isinstance(df, pd.DataFrame) assert df.empty assert set(df.columns) == set(df_collection._schema.columns.keys()) incite_item_factory(user=u1, item=item) df = item.from_mysql() + assert isinstance(df, pd.DataFrame) assert not df.empty assert set(df.columns) == set(df_collection._schema.columns.keys()) if df_collection.data_type == DFCollectionType.LEDGER: @@ -270,13 +273,13 @@ class TestDFCollectionItemMethod: def test_from_mysql_standard( self, - df_collection_data_type, - df_collection, + df_collection_data_type: DFCollectionType, + df_collection: DFCollection, offset: str, duration: timedelta, user_factory: Callable[..., User], - product: product: Product, - incite_item_factory, + product: Product, + incite_item_factory: Callable[..., None], delete_df_collection: Callable[..., None], ): @@ -293,7 +296,7 @@ class TestDFCollectionItemMethod: # We're using parametrize, so this If statement is just to # confirm other Item Types will always raise an assertion with pytest.raises(expected_exception=AssertionError) as cm: - res = item.from_mysql_standard() + _ = item.from_mysql_standard() assert ( "Can't call from_mysql_standard for Ledger DFCollectionItem" in str(cm.value) @@ -304,32 +307,34 @@ class TestDFCollectionItemMethod: # Unlike .from_mysql_ledger(), .from_mysql_standard() will return # back and empty df with the correct columns in place df = item.from_mysql_standard() + assert isinstance(df, pd.DataFrame) assert df.empty assert set(df.columns) == set(df_collection._schema.columns.keys()) incite_item_factory(user=u1, item=item) df = item.from_mysql_standard() + assert isinstance(df, pd.DataFrame) assert not df.empty assert set(df.columns) == set(df_collection._schema.columns.keys()) assert df.shape[0] > 0 def test_from_mysql_ledger( self, - df_collection, + df_collection: DFCollection, user: User, create_main_accounts: Callable[..., None], offset: str, duration: timedelta, thl_web_rw: PostgresConfig, - thl_lm, - df_collection_data_type, + thl_ledger_manager: ThlLedgerManager, + df_collection_data_type: DFCollectionType, user_factory: Callable[..., User], - product: product: Product, - client_no_amm, - incite_item_factory, + product: Product, + client_no_amm: DaskClient, + incite_item_factory: Callable[..., None], delete_df_collection: Callable[..., None], - mnt_filepath, + mnt_filepath: GRLDatasets, ): if df_collection.data_type != DFCollectionType.LEDGER: @@ -370,17 +375,17 @@ class TestDFCollectionItemMethod: def test_to_archive( self, - df_collection, + df_collection: DFCollection, user: User, offset: str, duration: timedelta, - df_collection_data_type, + df_collection_data_type: DFCollectionType, user_factory: Callable[..., User], - product: product: Product, - client_no_amm, - incite_item_factory, + product: Product, + client_no_amm: DaskClient, + incite_item_factory: Callable[..., None], delete_df_collection: Callable[..., None], - mnt_filepath, + mnt_filepath: GRLDatasets, ): if df_collection.data_type in unsupported_mock_types: @@ -407,17 +412,17 @@ class TestDFCollectionItemMethod: def test__to_archive( self, - df_collection_data_type, - df_collection, + df_collection_data_type: DFCollectionType, + df_collection: DFCollection, user_factory: Callable[..., User], - product: product: Product, + product: Product, offset: str, duration: timedelta, - client_no_amm, + client_no_amm: DaskClient, user: User, - incite_item_factory, + incite_item_factory: Callable[..., None], delete_df_collection: Callable[..., None], - mnt_filepath, + mnt_filepath: GRLDatasets, ): """We already have a test for the "non-private" version of this, which primarily just uses the respective Client to determine if @@ -480,19 +485,19 @@ class TestDFCollectionItemMethod: @pytest.mark.skip def test_to_archive_numbered_partial( - self, df_collection_data_type, df_collection, offset: str, duration: timedelta + self, ): pass @pytest.mark.skip def test_initial_load( - self, df_collection_data_type, df_collection, offset: str, duration: timedelta + self, ): pass @pytest.mark.skip def test_clear_corrupt_archive( - self, df_collection_data_type, df_collection, offset: str, duration: timedelta + self, ): pass @@ -505,34 +510,40 @@ class TestDFCollectionItemMethodBase: @pytest.mark.skip def test_path_exists( - self, df_collection_data_type, offset: str, duration: timedelta + self, ): pass @pytest.mark.skip def test_next_numbered_path( - self, df_collection_data_type, offset: str, duration: timedelta + self, ): pass @pytest.mark.skip def test_search_highest_numbered_path( - self, df_collection_data_type, offset: str, duration: timedelta + self, + df_collection_data_type: DFCollectionType, + offset: str, + duration: timedelta, ): pass @pytest.mark.skip def test_tmp_filename( - self, df_collection_data_type, offset: str, duration: timedelta + self, ): pass @pytest.mark.skip - def test_tmp_path(self, df_collection_data_type, offset: str, duration: timedelta): + def test_tmp_path( + self, + ): pass def test_is_empty( - self, df_collection_data_type, df_collection, offset: str, duration: timedelta + self, + df_collection: DFCollection, ): """ test_has_empty was merged into this because item.has_empty is @@ -549,7 +560,8 @@ class TestDFCollectionItemMethodBase: assert item.has_empty() def test_has_partial_archive( - self, df_collection_data_type, df_collection, offset: str, duration: timedelta + self, + df_collection: DFCollection, ): for item in df_collection.items: assert not item.has_partial_archive() @@ -557,7 +569,8 @@ class TestDFCollectionItemMethodBase: assert item.has_partial_archive() def test_has_archive( - self, df_collection_data_type, df_collection, offset: str, duration: timedelta + self, + df_collection: DFCollection, ): for item in df_collection.items: # (1) Originally, nothing exists... so let's just make a file and @@ -594,7 +607,8 @@ class TestDFCollectionItemMethodBase: assert item.has_archive(include_empty=True) def test_delete_archive( - self, df_collection_data_type, df_collection, offset: str, duration: timedelta + self, + df_collection: DFCollection, ): for item in df_collection.items: item: DFCollectionItem @@ -617,7 +631,8 @@ class TestDFCollectionItemMethodBase: assert not item.partial_path.exists() def test_should_archive( - self, df_collection_data_type, df_collection, offset: str, duration: timedelta + self, + df_collection: DFCollection, ): schema: DataFrameSchema = df_collection._schema aa = schema.metadata[ARCHIVE_AFTER] @@ -635,12 +650,13 @@ class TestDFCollectionItemMethodBase: @pytest.mark.skip def test_set_empty( - self, df_collection_data_type, df_collection, offset: str, duration: timedelta + self, ): pass def test_valid_archive( - self, df_collection_data_type, df_collection, offset: str, duration: timedelta + self, + df_collection: DFCollection, ): # Originally, nothing has been saved or anything.. so confirm it # always comes back as None @@ -664,18 +680,19 @@ class TestDFCollectionItemMethodBase: @pytest.mark.skip def test_validate_df( - self, df_collection_data_type, df_collection, offset: str, duration: timedelta + self, ): pass @pytest.mark.skip def test_from_archive( - self, df_collection_data_type, df_collection, offset: str, duration: timedelta + self, ): pass def test__to_dict( - self, df_collection_data_type, df_collection, offset: str, duration: timedelta + self, + df_collection: DFCollection, ): for item in df_collection.items: @@ -694,19 +711,19 @@ class TestDFCollectionItemMethodBase: @pytest.mark.skip def test_delete_partial( - self, df_collection_data_type, df_collection, offset: str, duration: timedelta + self, ): pass @pytest.mark.skip def test_cleanup_partials( - self, df_collection_data_type, df_collection, offset: str, duration: timedelta + self, ): pass @pytest.mark.skip def test_delete_dangling_partials( - self, df_collection_data_type, df_collection, offset: str, duration: timedelta + self, ): pass @@ -726,7 +743,9 @@ async def test_client(client, s, worker): ) @gen_cluster(client=True, nthreads=[("127.0.0.1", 1)]) @pytest.mark.anyio -async def test_client_parametrize(c, s, w, df_collection_data_type, offset: str): +async def test_client_parametrize( + c, s, w, df_collection_data_type: DFCollectionType, offset: str +): """c,s,a are all required - the secondary Worker (b) is not required""" assert isinstance(c, Client), f"c is not Client, it's {type(c)}" @@ -750,17 +769,12 @@ class TestDFCollectionItemFunctionalTest: def test_to_archive_and_ddf( self, - df_collection_data_type, - offset: str, - duration: timedelta, - client_no_amm, - df_collection, - user: User, + client_no_amm: DaskClient, + df_collection: DFCollection, user_factory: Callable[..., User], - product: product: Product, - incite_item_factory, + product: Product, + incite_item_factory: Callable[..., None], delete_df_collection: Callable[..., None], - mnt_filepath: GRLDatasets, ): if df_collection.data_type in unsupported_mock_types: @@ -799,17 +813,11 @@ class TestDFCollectionItemFunctionalTest: def test_filesize_estimate( self, - df_collection, - user: User, - offset: str, - duration: timedelta, - client_no_amm, + df_collection: DFCollection, user_factory: Callable[..., User], - product: product: Product, - df_collection_data_type, - incite_item_factory, + product: Product, + incite_item_factory: Callable[..., None], delete_df_collection: Callable[..., None], - mnt_filepath: GRLDatasets, ): """A functional test to write some Parquet files for the DFCollection and then confirm that the files get written @@ -846,16 +854,12 @@ class TestDFCollectionItemFunctionalTest: def test_to_archive_client( self, - client_no_amm, - df_collection, + client_no_amm: DaskClient, + df_collection: DFCollection, user_factory: Callable[..., User], - product: product: Product, - offset: str, - duration: timedelta, - df_collection_data_type, - incite_item_factory, + product: Product, + incite_item_factory: Callable[..., None], delete_df_collection: Callable[..., None], - mnt_filepath: GRLDatasets, ): delete_df_collection(coll=df_collection) @@ -885,7 +889,8 @@ class TestDFCollectionItemFunctionalTest: @pytest.mark.skip def test_get_items( - self, df_collection, product: product: Product, offset: str, duration: timedelta + self, + df_collection: DFCollection, ): with pytest.warns(expected_warning=ResourceWarning) as cm: df_collection.get_items_last365() @@ -898,16 +903,11 @@ class TestDFCollectionItemFunctionalTest: def test_saving_protections( self, - client_no_amm, - df_collection_data_type, - df_collection, - incite_item_factory, + df_collection: DFCollection, + incite_item_factory: Callable[..., None], delete_df_collection: Callable[..., None], user_factory: Callable[..., User], - product: product: Product, - offset: str, - duration: timedelta, - mnt_filepath: GRLDatasets, + product: Product, ): """Don't allow creating an archive for data that will likely be overwritten or updated @@ -939,15 +939,8 @@ class TestDFCollectionItemFunctionalTest: def test_empty_item( self, - client_no_amm, - df_collection_data_type, - df_collection, - incite_item_factory, + df_collection: DFCollection, delete_df_collection: Callable[..., None], - user: User, - offset: str, - duration: timedelta, - mnt_filepath: GRLDatasets, ): delete_df_collection(coll=df_collection) @@ -967,16 +960,12 @@ class TestDFCollectionItemFunctionalTest: def test_file_touching( self, - client_no_amm, - df_collection_data_type, - df_collection, - incite_item_factory, + client_no_amm: DaskClient, + df_collection: DFCollection, + incite_item_factory: Callable[..., None], delete_df_collection: Callable[..., None], user_factory: Callable[..., User], - product: product: Product, - offset: str, - duration: timedelta, - mnt_filepath, + product: Product, ): delete_df_collection(coll=df_collection) diff --git a/tests/incite/mergers/foundations/test_enriched_session.py b/tests/incite/mergers/foundations/test_enriched_session.py index 8254d81..2a161e4 100644 --- a/tests/incite/mergers/foundations/test_enriched_session.py +++ b/tests/incite/mergers/foundations/test_enriched_session.py @@ -1,3 +1,6 @@ +from __future__ import annotations + +from collections.abc import Callable from datetime import UTC, datetime, timedelta from decimal import Decimal from itertools import product @@ -5,10 +8,24 @@ from itertools import product import dask.dataframe as dd import pandas as pd import pytest +from dask.distributed import Client as DaskClient +from generalresearch.incite.collections.thl_web import ( + SessionDFCollection, + WallDFCollection, +) +from generalresearch.incite.mergers.foundations.enriched_session import ( + EnrichedSessionMerge, +) from generalresearch.incite.schemas.admin_responses import ( AdminPOPSessionSchema, ) +from generalresearch.models.admin.request import ( + ReportRequest, +) +from generalresearch.models.thl.product import Product +from generalresearch.models.thl.session import Session +from generalresearch.models.thl.user import User from generalresearch.pg_helper import PostgresConfig @@ -25,21 +42,20 @@ class TestEnrichedSession: def test_base( self, - client_no_amm, + client_no_amm: DaskClient, product: Product, user_factory: Callable[..., User], - wall_collection, - session_collection, - enriched_session_merge, + wall_collection: WallDFCollection, + session_collection: SessionDFCollection, + enriched_session_merge: EnrichedSessionMerge, thl_web_rr: PostgresConfig, delete_df_collection: Callable[..., None], - incite_item_factory, + incite_item_factory: Callable[..., None], ): - from generalresearch.models.thl.user import User delete_df_collection(coll=session_collection) - u1: User = user_factory(product=product: Product, created=session_collection.start) + u1: User = user_factory(product=product, created=session_collection.start) for item in session_collection.items: incite_item_factory(item=item, user=u1) @@ -52,7 +68,7 @@ class TestEnrichedSession: client=client_no_amm, wall_coll=wall_collection, session_coll=session_collection, - pg_config=thl_web_rr: PostgresConfig, + pg_config=thl_web_rr, ) # -- @@ -85,16 +101,16 @@ class TestEnrichedSessionAdmin: def test_to_admin_response( self, - event_report_request, - enriched_session_merge, - client_no_amm, - wall_collection, - session_collection, + event_report_request: ReportRequest, + enriched_session_merge: EnrichedSessionMerge, + client_no_amm: DaskClient, + wall_collection: WallDFCollection, + session_collection: SessionDFCollection, thl_web_rr: PostgresConfig, - session_report_request, + session_report_request: ReportRequest, user_factory: Callable[..., User], - start, - session_factory, + start: datetime, + session_factory: Callable[..., Session], product_factory: Callable[..., Product], delete_df_collection: Callable[..., None], ): @@ -107,7 +123,7 @@ class TestEnrichedSessionAdmin: for p in [p1, p2]: u = user_factory(product=p) for i in range(50): - s = session_factory( + _ = session_factory( user=u, wall_count=1, wall_req_cpi=Decimal("1.00"), @@ -120,7 +136,7 @@ class TestEnrichedSessionAdmin: client=client_no_amm, session_coll=session_collection, wall_coll=wall_collection, - pg_config=thl_web_rr: PostgresConfig, + pg_config=thl_web_rr, ) df = enriched_session_merge.to_admin_response( diff --git a/tests/incite/mergers/foundations/test_enriched_task_adjust.py b/tests/incite/mergers/foundations/test_enriched_task_adjust.py index a33a55a..0606b6f 100644 --- a/tests/incite/mergers/foundations/test_enriched_task_adjust.py +++ b/tests/incite/mergers/foundations/test_enriched_task_adjust.py @@ -1,9 +1,28 @@ +from __future__ import annotations + +from collections.abc import Callable from datetime import timedelta from itertools import product as iter_product import dask.dataframe as dd import pandas as pd import pytest +from dask.distributed import Client as DaskClient + +from generalresearch.incite.collections.thl_web import ( + SessionDFCollection, + TaskAdjustmentDFCollection, + WallDFCollection, +) +from generalresearch.incite.mergers.foundations.enriched_task_adjust import ( + EnrichedTaskAdjustMerge, +) +from generalresearch.incite.mergers.foundations.enriched_wall import ( + EnrichedWallMerge, +) +from generalresearch.models.thl.product import Product +from generalresearch.models.thl.user import User +from generalresearch.pg_helper import PostgresConfig @pytest.mark.parametrize( @@ -20,19 +39,18 @@ class TestEnrichedTaskAdjust: @pytest.mark.skip def test_base( self, - client_no_amm, + client_no_amm: DaskClient, user_factory: Callable[..., User], product: Product, - task_adj_collection, - wall_collection, - session_collection, - enriched_wall_merge, - enriched_task_adjust_merge, - incite_item_factory, + task_adj_collection: TaskAdjustmentDFCollection, + wall_collection: WallDFCollection, + session_collection: SessionDFCollection, + enriched_wall_merge: EnrichedWallMerge, + enriched_task_adjust_merge: EnrichedTaskAdjustMerge, + incite_item_factory: Callable[..., None], delete_df_collection: Callable[..., None], thl_web_rr: PostgresConfig, ): - from generalresearch.models.thl.user import User # -- Build & Setup delete_df_collection(coll=session_collection) @@ -48,14 +66,14 @@ class TestEnrichedTaskAdjust: client=client_no_amm, session_coll=session_collection, wall_coll=wall_collection, - pg_config=thl_web_rr: PostgresConfig, + pg_config=thl_web_rr, ) enriched_task_adjust_merge.build( client=client_no_amm, task_adjust_coll=task_adj_collection, enriched_wall=enriched_wall_merge, - pg_config=thl_web_rr: PostgresConfig, + pg_config=thl_web_rr, ) # -- diff --git a/tests/incite/mergers/foundations/test_enriched_wall.py b/tests/incite/mergers/foundations/test_enriched_wall.py index a0ca4dd..0cb8f60 100644 --- a/tests/incite/mergers/foundations/test_enriched_wall.py +++ b/tests/incite/mergers/foundations/test_enriched_wall.py @@ -1,3 +1,4 @@ +from collections.abc import Callable from datetime import UTC, datetime, timedelta from decimal import Decimal from itertools import product as iter_product @@ -5,11 +6,23 @@ from itertools import product as iter_product import dask.dataframe as dd import pandas as pd import pytest +from dask.distributed import Client as DaskClient + +from generalresearch.incite.collections.thl_web import ( + SessionDFCollection, + WallDFCollection, +) # noinspection PyUnresolvedReferences from generalresearch.incite.mergers.foundations.enriched_wall import ( + EnrichedWallMerge, EnrichedWallMergeItem, ) +from generalresearch.models.admin.request import ReportRequest +from generalresearch.models.thl.product import Product +from generalresearch.models.thl.session import Session +from generalresearch.models.thl.user import User +from generalresearch.pg_helper import PostgresConfig @pytest.mark.parametrize( @@ -20,22 +33,21 @@ class TestEnrichedWall: def test_base( self, - client_no_amm, + client_no_amm: DaskClient, product: Product, user_factory: Callable[..., User], - wall_collection, + wall_collection: WallDFCollection, thl_web_rr: PostgresConfig, - session_collection, - enriched_wall_merge, + session_collection: SessionDFCollection, + enriched_wall_merge: EnrichedWallMerge, delete_df_collection: Callable[..., None], - incite_item_factory, + incite_item_factory: Callable[..., None], ): - from generalresearch.models.thl.user import User # -- Build & Setup delete_df_collection(coll=session_collection) delete_df_collection(coll=wall_collection) - u1: User = user_factory(product=product: Product, created=session_collection.start) + u1: User = user_factory(product=product, created=session_collection.start) for item in session_collection.items: incite_item_factory(item=item, user=u1) @@ -48,7 +60,7 @@ class TestEnrichedWall: client=client_no_amm, wall_coll=wall_collection, session_coll=session_collection, - pg_config=thl_web_rr: PostgresConfig, + pg_config=thl_web_rr, ) # -- @@ -63,19 +75,19 @@ class TestEnrichedWall: def test_base_item( self, - client_no_amm, + client_no_amm: DaskClient, product: Product, user_factory: Callable[..., User], - wall_collection, - session_collection, - enriched_wall_merge, + wall_collection: WallDFCollection, + session_collection: SessionDFCollection, + enriched_wall_merge: EnrichedWallMerge, delete_df_collection: Callable[..., None], thl_web_rr: PostgresConfig, - incite_item_factory, + incite_item_factory: Callable[..., None], ): # -- Build & Setup delete_df_collection(coll=session_collection) - u = user_factory(product=product: Product, created=session_collection.start) + u = user_factory(product=product, created=session_collection.start) for item in session_collection.items: incite_item_factory(item=item, user=u) @@ -87,7 +99,7 @@ class TestEnrichedWall: client=client_no_amm, wall_coll=wall_collection, session_coll=session_collection, - pg_config=thl_web_rr: PostgresConfig, + pg_config=thl_web_rr, ) # -- @@ -99,14 +111,14 @@ class TestEnrichedWall: try: modified_time1 = path.stat().st_mtime - except Exception: + except OSError: modified_time1 = 0 item.build( client=client_no_amm, wall_coll=wall_collection, session_coll=session_collection, - pg_config=thl_web_rr: PostgresConfig, + pg_config=thl_web_rr, ) modified_time2 = path.stat().st_mtime @@ -150,7 +162,12 @@ class TestEnrichedWallToAdmin: def duration(self) -> timedelta | None: return timedelta(days=5) - def test_empty(self, enriched_wall_merge, client_no_amm, start): + def test_empty( + self, + enriched_wall_merge: EnrichedWallMerge, + client_no_amm: DaskClient, + start: datetime, + ): from generalresearch.models.admin.request import ReportRequest rr = ReportRequest.model_validate({"interval": "5min", "start": start}) @@ -167,18 +184,18 @@ class TestEnrichedWallToAdmin: def test_to_admin_response( self, - event_report_request, - enriched_wall_merge, - client_no_amm, - wall_collection, - session_collection, + event_report_request: ReportRequest, + enriched_wall_merge: EnrichedWallMerge, + client_no_amm: DaskClient, + wall_collection: WallDFCollection, + session_collection: SessionDFCollection, thl_web_rr: PostgresConfig, - user, - session_factory, + user: User, + session_factory: Callable[..., Session], delete_df_collection: Callable[..., None], product_factory: Callable[..., Product], user_factory: Callable[..., User], - start, + start: datetime, ): delete_df_collection(coll=wall_collection) delete_df_collection(coll=session_collection) @@ -189,7 +206,7 @@ class TestEnrichedWallToAdmin: for p in [p1, p2]: u = user_factory(product=p) for i in range(50): - s = session_factory( + _ = session_factory( user=u, wall_count=2, wall_req_cpi=Decimal("1.00"), @@ -203,7 +220,7 @@ class TestEnrichedWallToAdmin: client=client_no_amm, wall_coll=wall_collection, session_coll=session_collection, - pg_config=thl_web_rr: PostgresConfig, + pg_config=thl_web_rr, ) df = enriched_wall_merge.to_admin_response( diff --git a/tests/incite/mergers/test_merge_collection.py b/tests/incite/mergers/test_merge_collection.py index 15fa4db..cf8315f 100644 --- a/tests/incite/mergers/test_merge_collection.py +++ b/tests/incite/mergers/test_merge_collection.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from datetime import UTC, datetime, timedelta from itertools import product @@ -5,12 +7,13 @@ import pandas as pd import pytest from pandera.pandas import DataFrameSchema +from generalresearch.incite.base import GRLDatasets from generalresearch.incite.mergers import ( MergeCollection, MergeType, ) -merge_types = list(e for e in MergeType if e != MergeType.TEST) +merge_types = [e for e in MergeType if e != MergeType.TEST] @pytest.mark.parametrize( @@ -26,7 +29,11 @@ merge_types = list(e for e in MergeType if e != MergeType.TEST) ) class TestMergeCollection: - def test_init(self, mnt_filepath, merge_type, offset, duration, start): + def test_init( + self, + mnt_filepath: GRLDatasets, + merge_type: MergeType, + ): with pytest.raises(expected_exception=ValueError) as cm: MergeCollection(archive_path=mnt_filepath.data_src) assert "Must explicitly provide a merge_type" in str(cm.value) @@ -37,7 +44,14 @@ class TestMergeCollection: ) assert instance.merge_type == merge_type - def test_items(self, mnt_filepath, merge_type, offset, duration, start): + def test_items( + self, + mnt_filepath: GRLDatasets, + merge_type: MergeType, + offset: str, + duration: timedelta, + start: datetime, + ): instance = MergeCollection( merge_type=merge_type, offset=offset, @@ -48,7 +62,14 @@ class TestMergeCollection: assert len(instance.interval_range) == len(instance.items) - def test_progress(self, mnt_filepath, merge_type, offset, duration, start): + def test_progress( + self, + mnt_filepath: GRLDatasets, + merge_type: MergeType, + offset: str, + duration: timedelta, + start: datetime, + ): instance = MergeCollection( merge_type=merge_type, offset=offset, @@ -62,7 +83,11 @@ class TestMergeCollection: assert instance.progress.shape[1] == 7 assert instance.progress["group_by"].isnull().all() - def test_schema(self, mnt_filepath, merge_type, offset, duration, start): + def test_schema( + self, + mnt_filepath: GRLDatasets, + merge_type: MergeType, + ): instance = MergeCollection( merge_type=merge_type, archive_path=mnt_filepath.archive_path(enum_type=merge_type), @@ -70,7 +95,14 @@ class TestMergeCollection: assert isinstance(instance._schema, DataFrameSchema) - def test_load(self, mnt_filepath, merge_type, offset, duration, start): + def test_load( + self, + mnt_filepath: GRLDatasets, + merge_type: MergeType, + offset: str, + duration: timedelta, + start: datetime, + ): instance = MergeCollection( merge_type=merge_type, start=start, @@ -82,7 +114,14 @@ class TestMergeCollection: # Confirm that there are no archives available yet assert instance.progress.has_archive.eq(False).all() - def test_get_items(self, mnt_filepath, merge_type, offset, duration, start): + def test_get_items( + self, + mnt_filepath: GRLDatasets, + merge_type: MergeType, + offset: str, + duration: timedelta, + start: datetime, + ): instance = MergeCollection( start=start, finished=start + duration, diff --git a/tests/incite/mergers/test_merge_collection_item.py b/tests/incite/mergers/test_merge_collection_item.py index 3d0b644..5ca2f6b 100644 --- a/tests/incite/mergers/test_merge_collection_item.py +++ b/tests/incite/mergers/test_merge_collection_item.py @@ -1,10 +1,16 @@ +from __future__ import annotations + from datetime import timedelta from itertools import product from pathlib import PurePath import pytest -from generalresearch.incite.mergers import MergeCollectionItem, MergeType +from generalresearch.incite.mergers import ( + MergeCollection, + MergeCollectionItem, + MergeType, +) @pytest.mark.parametrize( @@ -19,7 +25,10 @@ from generalresearch.incite.mergers import MergeCollectionItem, MergeType ) class TestMergeCollectionItem: - def test_file_naming(self, merge_collection, offset, duration, start): + def test_file_naming( + self, + merge_collection: MergeCollection, + ): assert len(merge_collection.items) == 25 items: list[MergeCollectionItem] = merge_collection.items @@ -34,7 +43,10 @@ class TestMergeCollectionItem: assert i._collection.offset in i.filename assert i.start.strftime("%Y-%m-%d-%H-%M-%S") in i.filename - def test_archives(self, merge_collection, offset, duration, start): + def test_archives( + self, + merge_collection: MergeCollection, + ): assert len(merge_collection.items) == 25 for i in merge_collection.items: @@ -44,10 +56,13 @@ class TestMergeCollectionItem: assert not i.has_partial_archive() assert i.has_archive() == i.path_exists(generic_path=i.path) - res = set([i.should_archive() for i in merge_collection.items]) + res = {i.should_archive() for i in merge_collection.items} assert len(res) == 1 - def test_item_to_archive(self, merge_collection, offset, duration, start): + def test_item_to_archive( + self, + merge_collection: MergeCollection, + ): for item in merge_collection.items: item: MergeCollectionItem assert not item.has_archive() diff --git a/tests/incite/mergers/test_pop_ledger.py b/tests/incite/mergers/test_pop_ledger.py index d054eb6..529a641 100644 --- a/tests/incite/mergers/test_pop_ledger.py +++ b/tests/incite/mergers/test_pop_ledger.py @@ -1,12 +1,25 @@ +from __future__ import annotations + +from collections.abc import Callable from datetime import UTC, datetime, timedelta from itertools import product as iter_product import pandas as pd import pytest +from dask.distributed import Client as DaskClient +from generalresearch.incite.base import GRLDatasets +from generalresearch.incite.collections.thl_web import ( + LedgerDFCollection, + SessionDFCollection, +) +from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge from generalresearch.incite.schemas.mergers.pop_ledger import ( numerical_col_names, ) +from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager +from generalresearch.models.thl.product import Product +from generalresearch.models.thl.user import User @pytest.mark.parametrize( @@ -30,20 +43,20 @@ class TestMergePOPLedger: def test_base( self, - client_no_amm, - ledger_collection, - pop_ledger_merge, + client_no_amm: DaskClient, + ledger_collection: LedgerDFCollection, + pop_ledger_merge: PopLedgerMerge, product: Product, user_factory: Callable[..., User], create_main_accounts: Callable[..., None], - thl_lm, + thl_ledger_manager: ThlLedgerManager, delete_df_collection: Callable[..., None], - incite_item_factory, + incite_item_factory: Callable[..., None], delete_ledger_db: Callable[..., None], ): from generalresearch.models.thl.ledger import LedgerAccount - u = user_factory(product=product: Product, created=ledger_collection.start) + u = user_factory(product=product, created=ledger_collection.start) # -- Build & Setup delete_ledger_db() @@ -73,19 +86,21 @@ class TestMergePOPLedger: # -- - user_wallet_account: LedgerAccount = thl_lm.get_account_or_create_user_wallet( - user=u + user_wallet_account: LedgerAccount = ( + thl_ledger_manager.get_account_or_create_user_wallet(user=u) + ) + cash_account: LedgerAccount = thl_ledger_manager.get_account_cash() + rev_account: LedgerAccount = ( + thl_ledger_manager.get_account_task_complete_revenue() ) - cash_account: LedgerAccount = thl_lm.get_account_cash() - rev_account: LedgerAccount = thl_lm.get_account_task_complete_revenue() item_finishes = [i.finish for i in ledger_collection.items] item_finishes.sort(reverse=True) last_item_finish = item_finishes[0] # Pure SQL based lookups - cash_balance: int = thl_lm.get_account_balance(account=cash_account) - rev_balance: int = thl_lm.get_account_balance(account=rev_account) + cash_balance: int = thl_ledger_manager.get_account_balance(account=cash_account) + rev_balance: int = thl_ledger_manager.get_account_balance(account=rev_account) assert cash_balance > rev_balance # (1) Test Cash Account @@ -123,39 +138,42 @@ class TestMergePOPLedger: def test_pydantic_init( self, - client_no_amm, - ledger_collection, - pop_ledger_merge, - mnt_filepath, + client_no_amm: DaskClient, + ledger_collection: LedgerDFCollection, + pop_ledger_merge: PopLedgerMerge, + mnt_filepath: GRLDatasets, product: Product, user_factory: Callable[..., User], create_main_accounts: Callable[..., None], - offset, - duration, - start, - thl_lm, - incite_item_factory, + offset: str, + duration: timedelta, + start: datetime, + thl_ledger_manager: ThlLedgerManager, + incite_item_factory: Callable[..., None], delete_df_collection: Callable[..., None], delete_ledger_db: Callable[..., None], - session_collection, + session_collection: SessionDFCollection, ): from generalresearch.models.thl.finance import ProductBalances from generalresearch.models.thl.ledger import LedgerAccount from generalresearch.models.thl.product import Product - u = user_factory(product=product: Product, created=session_collection.start) + u = user_factory(product=product, created=session_collection.start) assert ledger_collection.finished is not None - assert isinstance(u.product: Product, Product) + assert isinstance(u.product, Product) delete_ledger_db() - create_main_accounts(), + create_main_accounts() + delete_df_collection(coll=ledger_collection) - bp_account: LedgerAccount = thl_lm.get_account_or_create_bp_wallet( + bp_account: LedgerAccount = thl_ledger_manager.get_account_or_create_bp_wallet( product=u.product ) - cash_account: LedgerAccount = thl_lm.get_account_cash() - rev_account: LedgerAccount = thl_lm.get_account_task_complete_revenue() + cash_account: LedgerAccount = thl_ledger_manager.get_account_cash() + rev_account: LedgerAccount = ( + thl_ledger_manager.get_account_task_complete_revenue() + ) for item in ledger_collection.items: incite_item_factory(item=item, user=u) @@ -185,8 +203,10 @@ class TestMergePOPLedger: assert instance.payout == instance.net == instance.bp_payment_credit assert instance.available_balance < instance.net assert instance.available_balance + instance.retainer == instance.net - assert instance.balance == thl_lm.get_account_balance(bp_account) - assert df["bp_payment.CREDIT"].sum() == thl_lm.get_account_balance(bp_account) + assert instance.balance == thl_ledger_manager.get_account_balance(bp_account) + assert df["bp_payment.CREDIT"].sum() == thl_ledger_manager.get_account_balance( + bp_account + ) # (2) Filter by the Cash Account ddf = pop_ledger_merge.ddf( @@ -199,7 +219,7 @@ class TestMergePOPLedger: ) df: pd.DataFrame = client_no_amm.compute(collections=ddf, sync=True) - cash_balance: int = thl_lm.get_account_balance(account=cash_account) + cash_balance: int = thl_ledger_manager.get_account_balance(account=cash_account) assert df["bp_payment.CREDIT"].sum() == 0 assert cash_balance > 0 assert df["mp_payment.CREDIT"].sum() == 0 @@ -216,7 +236,7 @@ class TestMergePOPLedger: ) df: pd.DataFrame = client_no_amm.compute(collections=ddf, sync=True) - rev_balance: int = thl_lm.get_account_balance(account=rev_account) + rev_balance: int = thl_ledger_manager.get_account_balance(account=rev_account) assert rev_balance == 0 assert df["bp_payment.CREDIT"].sum() == 0 assert df["mp_payment.DEBIT"].sum() == 0 @@ -224,27 +244,28 @@ class TestMergePOPLedger: def test_resample( self, - client_no_amm, - ledger_collection, - pop_ledger_merge, - mnt_filepath, + client_no_amm: DaskClient, + ledger_collection: LedgerDFCollection, + pop_ledger_merge: PopLedgerMerge, + mnt_filepath: GRLDatasets, user_factory: Callable[..., User], product: Product, create_main_accounts: Callable[..., None], - offset, - duration, - start, - thl_lm, + offset: str, + duration: timedelta, + start: datetime, + thl_ledger_manager: ThlLedgerManager, delete_df_collection: Callable[..., None], - incite_item_factory, + incite_item_factory: Callable[..., None], ): - from generalresearch.models.thl.user import User assert ledger_collection.finished is not None delete_df_collection(coll=ledger_collection) u1: User = user_factory(product=product) - bp_account = thl_lm.get_account_or_create_bp_wallet(product=u1.product) + bp_account = thl_ledger_manager.get_account_or_create_bp_wallet( + product=u1.product + ) for item in ledger_collection.items: incite_item_factory(user=u1, item=item) @@ -274,7 +295,7 @@ class TestMergePOPLedger: assert isinstance(df.index, pd.Index) assert isinstance(df.index, pd.DatetimeIndex) - bp_account_balance = thl_lm.get_account_balance(account=bp_account) + bp_account_balance = thl_ledger_manager.get_account_balance(account=bp_account) # Initial sum initial_sum = df.sum().sum() diff --git a/tests/incite/mergers/test_ym_survey_merge.py b/tests/incite/mergers/test_ym_survey_merge.py index a0b8b87..8a4897b 100644 --- a/tests/incite/mergers/test_ym_survey_merge.py +++ b/tests/incite/mergers/test_ym_survey_merge.py @@ -1,8 +1,24 @@ +from __future__ import annotations + +from collections.abc import Callable from datetime import UTC, datetime, timedelta from itertools import product import pandas as pd import pytest +from dask.distributed import Client as DaskClient + +from generalresearch.incite.collections.thl_web import ( + SessionDFCollection, + WallDFCollection, +) +from generalresearch.incite.mergers.foundations.enriched_session import ( + EnrichedSessionMerge, +) +from generalresearch.incite.mergers.ym_survey_wall import YMSurveyWallMerge +from generalresearch.models.thl.product import Product +from generalresearch.models.thl.user import User +from generalresearch.pg_helper import PostgresConfig # noinspection PyUnresolvedReferences @@ -27,21 +43,20 @@ class TestYMSurveyMerge: def test_base( self, - client_no_amm, + client_no_amm: DaskClient, user_factory: Callable[..., User], product: Product, - ym_survey_wall_merge, - wall_collection, - session_collection, - enriched_session_merge, + ym_survey_wall_merge: YMSurveyWallMerge, + wall_collection: WallDFCollection, + session_collection: SessionDFCollection, + enriched_session_merge: EnrichedSessionMerge, delete_df_collection: Callable[..., None], - incite_item_factory, + incite_item_factory: Callable[..., None], thl_web_rr: PostgresConfig, ): - from generalresearch.models.thl.user import User delete_df_collection(coll=session_collection) - user: User = user_factory(product=product: Product, created=session_collection.start) + user: User = user_factory(product=product, created=session_collection.start) # -- Build & Setup assert ym_survey_wall_merge.start is None @@ -61,15 +76,15 @@ class TestYMSurveyMerge: client=client_no_amm, session_coll=session_collection, wall_coll=wall_collection, - pg_config=thl_web_rr: PostgresConfig, + pg_config=thl_web_rr, ) assert enriched_session_merge.progress.has_archive.eq(True).all() ddf = enriched_session_merge.ddf() - df: pd.DataFrame = client_no_amm.compute(collections=ddf, sync=True) + df1: pd.DataFrame | None = client_no_amm.compute(collections=ddf, sync=True) - assert isinstance(df, pd.DataFrame) - assert not df.empty + assert isinstance(df1, pd.DataFrame) + assert not df1.empty # -- @@ -83,18 +98,18 @@ class TestYMSurveyMerge: # -- ddf = ym_survey_wall_merge.ddf() - df: pd.DataFrame = client_no_amm.compute(collections=ddf, sync=True) + df2: pd.DataFrame | None = client_no_amm.compute(collections=ddf, sync=True) - assert isinstance(df, pd.DataFrame) - assert not df.empty + assert isinstance(df2, pd.DataFrame) + assert not df2.empty # -- - assert df.product_id.nunique() == 1 - assert df.team_id.nunique() == 1 - assert df.source.nunique() > 1 + assert df2.product_id.nunique() == 1 + assert df2.team_id.nunique() == 1 + assert df2.source.nunique() > 1 - started_min_ts = df.started.min() - started_max_ts = df.started.max() + started_min_ts = df2.started.min() + started_max_ts = df2.started.max() assert type(started_min_ts) is pd.Timestamp assert type(started_max_ts) is pd.Timestamp diff --git a/tests/incite/schemas/test_admin_responses.py b/tests/incite/schemas/test_admin_responses.py index e98eecd..d2658ea 100644 --- a/tests/incite/schemas/test_admin_responses.py +++ b/tests/incite/schemas/test_admin_responses.py @@ -1,8 +1,11 @@ +from __future__ import annotations + from datetime import UTC, datetime, timedelta from random import sample import numpy as np import pandas as pd +import pandera as pa import pytest from generalresearch.incite.schemas import empty_dataframe_from_schema @@ -16,12 +19,14 @@ from generalresearch.locales import Localelator class TestAdminPOPSchema: schema_df = empty_dataframe_from_schema(AdminPOPSchema) countries = list(Localelator().get_all_countries())[:5] - dates = [datetime(year=2024, month=1, day=i, tzinfo=None) for i in range(1, 10)] + dates = [ + datetime(year=2024, month=1, day=i, tzinfo=None) for i in range(1, 10) # noqa + ] @classmethod def assign_valid_vals(cls, df: pd.DataFrame) -> pd.DataFrame: for c in df.columns: - check_attrs: dict = AdminPOPSchema.columns[c].checks[0].statistics + check_attrs = AdminPOPSchema.columns[c].checks[0].statistics df[c] = np.random.randint( check_attrs["min_value"], check_attrs["max_value"], df.shape[0] ) @@ -29,7 +34,7 @@ class TestAdminPOPSchema: return df def test_empty(self): - with pytest.raises(Exception): + with pytest.raises(pa.errors.SchemaError): AdminPOPSchema.validate(pd.DataFrame()) def test_new_empty_df(self): @@ -42,7 +47,7 @@ class TestAdminPOPSchema: def test_valid(self): # (1) Works with raw naive datetime dates = [ - datetime(year=2024, month=1, day=i, tzinfo=None).isoformat() + datetime(year=2024, month=1, day=i, tzinfo=None).isoformat() # noqa for i in range(1, 10) ] df = pd.DataFrame( @@ -57,7 +62,10 @@ class TestAdminPOPSchema: assert isinstance(df, pd.DataFrame) # (2) Works with isoformat naive datetime - dates = [datetime(year=2024, month=1, day=i, tzinfo=None) for i in range(1, 10)] + dates = [ + datetime(year=2024, month=1, day=i, tzinfo=None) # noqa + for i in range(1, 10) + ] df = pd.DataFrame( index=pd.MultiIndex.from_product( iterables=[dates, self.countries], names=["index0", "index1"] @@ -84,12 +92,12 @@ class TestAdminPOPSchema: # Initially, they're all set with a timezone timestmaps: list[pd.Timestamp] = [i for i in df.index.get_level_values(0)] - assert all([ts.tz == UTC for ts in timestmaps]) + assert all(ts.tz == UTC for ts in timestmaps) # After validation, the timezone is removed df = AdminPOPSchema.validate(df) timestmaps: list[pd.Timestamp] = [i for i in df.index.get_level_values(0)] - assert all([ts.tz is None for ts in timestmaps]) + assert all(ts.tz is None for ts in timestmaps) def test_index_tz_no_future_beyond_one_year(self): now = datetime.now(tz=UTC) @@ -123,12 +131,12 @@ class TestAdminPOPSchema: df = self.assign_valid_vals(df) vals = [i for i in df.index.get_level_values(1)] - assert all([isinstance(v, float) for v in vals]) + assert all(isinstance(v, float) for v in vals) df = AdminPOPSchema.validate(df, lazy=True) vals = [i for i in df.index.get_level_values(1)] - assert all([isinstance(v, str) for v in vals]) + assert all(isinstance(v, str) for v in vals) # --- int to str --- @@ -142,12 +150,12 @@ class TestAdminPOPSchema: df = self.assign_valid_vals(df) vals = [i for i in df.index.get_level_values(1)] - assert all([isinstance(v, int) for v in vals]) + assert all(isinstance(v, int) for v in vals) df = AdminPOPSchema.validate(df, lazy=True) vals = [i for i in df.index.get_level_values(1)] - assert all([isinstance(v, str) for v in vals]) + assert all(isinstance(v, str) for v in vals) # a = 1 assert isinstance(df, pd.DataFrame) @@ -170,7 +178,7 @@ class TestAdminPOPSchema: assert isinstance(df, pd.DataFrame) timestmaps: list[pd.Timestamp] = [i for i in df.index.get_level_values(0)] - assert all([ts.tz is None for ts in timestmaps]) + assert all(ts.tz is None for ts in timestmaps) # (2) Timezones are removed dates = [ @@ -187,12 +195,12 @@ class TestAdminPOPSchema: # Has tz before validation, and none after timestmaps: list[pd.Timestamp] = [i for i in df.index.get_level_values(0)] - assert all([ts.tz is UTC for ts in timestmaps]) + assert all(ts.tz is UTC for ts in timestmaps) df = AdminPOPSchema.validate(df, lazy=True) timestmaps: list[pd.Timestamp] = [i for i in df.index.get_level_values(0)] - assert all([ts.tz is None for ts in timestmaps]) + assert all(ts.tz is None for ts in timestmaps) def test_clipping(self): df = pd.DataFrame( diff --git a/tests/incite/schemas/test_thl_web.py b/tests/incite/schemas/test_thl_web.py index 7f4434b..9b34ce0 100644 --- a/tests/incite/schemas/test_thl_web.py +++ b/tests/incite/schemas/test_thl_web.py @@ -16,7 +16,7 @@ class TestWallSchema: df = pd.DataFrame(columns=THLWallSchema.columns.keys()) - with pytest.raises(SchemaError) as cm: + with pytest.raises(SchemaError): THLWallSchema.validate(df) def test_no_rows(self): @@ -24,7 +24,7 @@ class TestWallSchema: df = pd.DataFrame(index=["uuid"], columns=THLWallSchema.columns.keys()) - with pytest.raises(SchemaError) as cm: + with pytest.raises(SchemaError): THLWallSchema.validate(df) def test_new_empty_df(self): @@ -50,7 +50,7 @@ class TestSessionSchema: df = pd.DataFrame(columns=THLSessionSchema.columns.keys()) df.set_index("uuid", inplace=True) - with pytest.raises(SchemaError) as cm: + with pytest.raises(SchemaError): THLSessionSchema.validate(df) def test_no_rows(self): @@ -58,7 +58,7 @@ class TestSessionSchema: df = pd.DataFrame(index=["id"], columns=THLSessionSchema.columns.keys()) - with pytest.raises(SchemaError) as cm: + with pytest.raises(SchemaError): THLSessionSchema.validate(df) def test_new_empty_df(self): diff --git a/tests/incite/test_collection_base.py b/tests/incite/test_collection_base.py index 7e1577a..d6ce2b1 100644 --- a/tests/incite/test_collection_base.py +++ b/tests/incite/test_collection_base.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from datetime import UTC, datetime, timedelta, timezone from os.path import exists as pexists from os.path import join as pjoin @@ -9,7 +11,7 @@ import pandas as pd import pytest from _pytest._code.code import ExceptionInfo -from generalresearch.incite.base import CollectionBase +from generalresearch.incite.base import CollectionBase, GRLDatasets AGO_15min = (datetime.now(tz=UTC) - timedelta(minutes=15)).replace(microsecond=0) AGO_1HR = (datetime.now(tz=UTC) - timedelta(hours=1)).replace(microsecond=0) @@ -17,11 +19,11 @@ AGO_2HR = (datetime.now(tz=UTC) - timedelta(hours=2)).replace(microsecond=0) class TestCollectionBase: - def test_init(self, mnt_filepath): + def test_init(self, mnt_filepath: GRLDatasets): instance = CollectionBase(archive_path=mnt_filepath.data_src) assert instance.df.empty is True - def test_init_df(self, mnt_filepath): + def test_init_df(self, mnt_filepath: GRLDatasets): # Only an empty pd.DataFrame can ever be provided instance = CollectionBase( df=pd.DataFrame({}), archive_path=mnt_filepath.data_src @@ -43,7 +45,7 @@ class TestCollectionBase: ) assert "Do not provide a pd.DataFrame" in str(cm.value) - def test_init_start(self, mnt_filepath): + def test_init_start(self, mnt_filepath: GRLDatasets): with pytest.raises(expected_exception=ValueError) as cm: cm: ExceptionInfo CollectionBase( @@ -74,7 +76,7 @@ class TestCollectionBase: cm.value ) - def test_init_archive_path(self, mnt_filepath): + def test_init_archive_path(self, mnt_filepath: GRLDatasets): """DirectoryPath is apparently smart enough to confirm that the directory path exists. """ @@ -99,7 +101,7 @@ class TestCollectionBase: CollectionBase(archive_path=new_path) assert "Path does not point to a directory" in str(cm.value) - def test_init_offset(self, mnt_filepath): + def test_init_offset(self, mnt_filepath: GRLDatasets): with pytest.raises(expected_exception=ValueError) as cm: cm: ExceptionInfo CollectionBase(offset="1:X", archive_path=mnt_filepath.data_src) @@ -118,14 +120,14 @@ class TestCollectionBase: class TestCollectionBaseProperties: - def test_items(self, mnt_filepath): + def test_items(self, mnt_filepath: GRLDatasets): with pytest.raises(expected_exception=NotImplementedError) as cm: cm: ExceptionInfo instance = CollectionBase(archive_path=mnt_filepath.data_src) - x = instance.items + _ = instance.items assert "Must override" in str(cm.value) - def test_interval_range(self, mnt_filepath): + def test_interval_range(self, mnt_filepath: GRLDatasets): instance = CollectionBase(archive_path=mnt_filepath.data_src) # Private method requires the end parameter with pytest.raises(expected_exception=AssertionError) as cm: @@ -147,7 +149,7 @@ class TestCollectionBaseProperties: assert res.is_monotonic_increasing assert res.is_unique - def test_interval_range2(self, mnt_filepath): + def test_interval_range2(self, mnt_filepath: GRLDatasets): instance = CollectionBase(archive_path=mnt_filepath.data_src) assert isinstance(instance.interval_range, list) @@ -166,16 +168,16 @@ class TestCollectionBaseProperties: ) assert len(instance.interval_range) == 2 - def test_progress(self, mnt_filepath): + def test_progress(self, mnt_filepath: GRLDatasets): with pytest.raises(expected_exception=NotImplementedError) as cm: cm: ExceptionInfo instance = CollectionBase( start=AGO_15min, offset="3min", archive_path=mnt_filepath.data_src ) - x = instance.progress + _ = instance.progress assert "Must override" in str(cm.value) - def test_progress2(self, mnt_filepath): + def test_progress2(self, mnt_filepath: GRLDatasets): instance = CollectionBase( start=AGO_2HR, offset="15min", @@ -184,10 +186,10 @@ class TestCollectionBaseProperties: assert instance.df.empty with pytest.raises(expected_exception=NotImplementedError) as cm: - df = instance.progress + _ = instance.progress assert "Must override" in str(cm.value) - def test_items2(self, mnt_filepath): + def test_items2(self, mnt_filepath: GRLDatasets): """There can't be a test for this because the Items need a path whic isn't possible in the generic form """ @@ -197,7 +199,7 @@ class TestCollectionBaseProperties: with pytest.raises(expected_exception=NotImplementedError) as cm: cm: ExceptionInfo - items = instance.items + _ = instance.items assert "Must override" in str(cm.value) # item = items[-3] @@ -208,19 +210,19 @@ class TestCollectionBaseProperties: # assert str(df.product_id.dtype) == "object" # assert str(ddf.product_id.dtype) == "string" - def test_items3(self, mnt_filepath): + def test_items3(self, mnt_filepath: GRLDatasets): instance = CollectionBase( start=AGO_2HR, offset="15min", archive_path=mnt_filepath.data_src, ) with pytest.raises(expected_exception=NotImplementedError) as cm: - item = instance.items[0] + _ = instance.items[0] assert "Must override" in str(cm.value) class TestCollectionBaseMethodsCleanup: - def test_fetch_force_rr_latest(self, mnt_filepath): + def test_fetch_force_rr_latest(self, mnt_filepath: GRLDatasets): coll = CollectionBase(archive_path=mnt_filepath.data_src) with pytest.raises(expected_exception=Exception) as cm: @@ -228,7 +230,7 @@ class TestCollectionBaseMethodsCleanup: coll.fetch_force_rr_latest(sources=[]) assert "Must override" in str(cm.value) - def test_fetch_all_paths(self, mnt_filepath): + def test_fetch_all_paths(self, mnt_filepath: GRLDatasets): coll = CollectionBase(archive_path=mnt_filepath.data_src) with pytest.raises(expected_exception=NotImplementedError) as cm: @@ -242,16 +244,16 @@ class TestCollectionBaseMethodsCleanup: class TestCollectionBaseMethodsCleanup: @pytest.mark.skip - def test_cleanup_partials(self, mnt_filepath): + def test_cleanup_partials(self, mnt_filepath: GRLDatasets): instance = CollectionBase(archive_path=mnt_filepath.data_src) assert instance.cleanup_partials() is None # it doesn't return anything - def test_clear_tmp_archives(self, mnt_filepath): + def test_clear_tmp_archives(self, mnt_filepath: GRLDatasets): instance = CollectionBase(archive_path=mnt_filepath.data_src) assert instance.clear_tmp_archives() is None # it doesn't return anything @pytest.mark.skip - def test_clear_corrupt_archives(self, mnt_filepath): + def test_clear_corrupt_archives(self, mnt_filepath: GRLDatasets): """TODO: expand this so it actually has corrupt archives that we check to see if they're removed """ @@ -259,14 +261,14 @@ class TestCollectionBaseMethodsCleanup: assert instance.clear_corrupt_archives() is None # it doesn't return anything @pytest.mark.skip - def test_rebuild_symlinks(self, mnt_filepath): + def test_rebuild_symlinks(self, mnt_filepath: GRLDatasets): instance = CollectionBase(archive_path=mnt_filepath.data_src) assert instance.rebuild_symlinks() is None class TestCollectionBaseMethodsSourceTiming: - def test_get_item(self, mnt_filepath): + def test_get_item(self, mnt_filepath: GRLDatasets): instance = CollectionBase(archive_path=mnt_filepath.data_src) i = pd.Interval(left=1, right=2, closed="left") @@ -274,7 +276,7 @@ class TestCollectionBaseMethodsSourceTiming: instance.get_item(interval=i) assert "Must override" in str(cm.value) - def test_get_item_start(self, mnt_filepath): + def test_get_item_start(self, mnt_filepath: GRLDatasets): instance = CollectionBase(archive_path=mnt_filepath.data_src) dt = datetime.now(tz=UTC) @@ -284,7 +286,7 @@ class TestCollectionBaseMethodsSourceTiming: instance.get_item_start(start=start) assert "Must override" in str(cm.value) - def test_get_items(self, mnt_filepath): + def test_get_items(self, mnt_filepath: GRLDatasets): instance = CollectionBase(archive_path=mnt_filepath.data_src) dt = datetime.now(tz=UTC) @@ -293,21 +295,21 @@ class TestCollectionBaseMethodsSourceTiming: instance.get_items(since=dt) assert "Must override" in str(cm.value) - def test_get_items_from_year(self, mnt_filepath): + def test_get_items_from_year(self, mnt_filepath: GRLDatasets): instance = CollectionBase(archive_path=mnt_filepath.data_src) with pytest.raises(expected_exception=NotImplementedError) as cm: instance.get_items_from_year(year=2020) assert "Must override" in str(cm.value) - def test_get_items_last90(self, mnt_filepath): + def test_get_items_last90(self, mnt_filepath: GRLDatasets): instance = CollectionBase(archive_path=mnt_filepath.data_src) with pytest.raises(expected_exception=NotImplementedError) as cm: instance.get_items_last90() assert "Must override" in str(cm.value) - def test_get_items_last365(self, mnt_filepath): + def test_get_items_last365(self, mnt_filepath: GRLDatasets): instance = CollectionBase(archive_path=mnt_filepath.data_src) with pytest.raises(expected_exception=NotImplementedError) as cm: diff --git a/tests/incite/test_collection_base_item.py b/tests/incite/test_collection_base_item.py index 7a0a581..e09f54a 100644 --- a/tests/incite/test_collection_base_item.py +++ b/tests/incite/test_collection_base_item.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from datetime import UTC, datetime from os.path import join as pjoin from pathlib import Path @@ -8,7 +10,7 @@ import pandas as pd import pytest from pydantic import ValidationError -from generalresearch.incite.base import CollectionItemBase +from generalresearch.incite.base import CollectionItemBase, GRLDatasets class TestCollectionItemBase: @@ -40,20 +42,20 @@ class TestCollectionItemBaseProperties: def test_finish(self): instance = CollectionItemBase() - with pytest.raises(expected_exception=AttributeError) as cm: - res = instance.finish + with pytest.raises(expected_exception=AttributeError): + _ = instance.finish def test_interval(self): instance = CollectionItemBase() - with pytest.raises(expected_exception=AttributeError) as cm: - res = instance.interval + with pytest.raises(expected_exception=AttributeError): + _ = instance.interval def test_filename(self): instance = CollectionItemBase() with pytest.raises(expected_exception=NotImplementedError) as cm: - res = instance.filename + _ = instance.filename assert "Do not use CollectionItemBase directly" in str(cm.value) @@ -61,7 +63,7 @@ class TestCollectionItemBaseProperties: instance = CollectionItemBase() with pytest.raises(expected_exception=NotImplementedError) as cm: - res = instance.filename + _ = instance.filename assert "Do not use CollectionItemBase directly" in str(cm.value) @@ -69,27 +71,27 @@ class TestCollectionItemBaseProperties: instance = CollectionItemBase() with pytest.raises(expected_exception=NotImplementedError) as cm: - res = instance.filename + _ = instance.filename assert "Do not use CollectionItemBase directly" in str(cm.value) def test_path(self): instance = CollectionItemBase() - with pytest.raises(expected_exception=AttributeError) as cm: - res = instance.path + with pytest.raises(expected_exception=AttributeError): + _ = instance.path def test_partial_path(self): instance = CollectionItemBase() - with pytest.raises(expected_exception=AttributeError) as cm: - res = instance.partial_path + with pytest.raises(expected_exception=AttributeError): + _ = instance.partial_path def test_empty_path(self): instance = CollectionItemBase() - with pytest.raises(expected_exception=AttributeError) as cm: - res = instance.empty_path + with pytest.raises(expected_exception=AttributeError): + _ = instance.empty_path class TestCollectionItemBaseMethods: @@ -106,41 +108,41 @@ class TestCollectionItemBaseMethods: instance = CollectionItemBase() with pytest.raises(expected_exception=NotImplementedError) as cm: - res = instance.tmp_filename() + _ = instance.tmp_filename() assert "Do not use CollectionItemBase directly" in str(cm.value) def test_tmp_path(self): instance = CollectionItemBase() - with pytest.raises(expected_exception=AttributeError) as cm: - res = instance.tmp_path() + with pytest.raises(expected_exception=AttributeError): + instance.tmp_path() def test_is_empty(self): instance = CollectionItemBase() - with pytest.raises(expected_exception=AttributeError) as cm: - res = instance.is_empty() + with pytest.raises(expected_exception=AttributeError): + instance.is_empty() def test_has_empty(self): instance = CollectionItemBase() - with pytest.raises(expected_exception=AttributeError) as cm: - res = instance.has_empty() + with pytest.raises(expected_exception=AttributeError): + instance.has_empty() def test_has_partial_archive(self): instance = CollectionItemBase() - with pytest.raises(expected_exception=AttributeError) as cm: - res = instance.has_partial_archive() + with pytest.raises(expected_exception=AttributeError): + instance.has_partial_archive() @pytest.mark.parametrize("include_empty", [True, False]) - def test_has_archive(self, include_empty): + def test_has_archive(self, include_empty: bool): instance = CollectionItemBase() - with pytest.raises(expected_exception=AttributeError) as cm: - res = instance.has_archive(include_empty=include_empty) + with pytest.raises(expected_exception=AttributeError): + instance.has_archive(include_empty=include_empty) - def test_delete_archive_file(self, mnt_filepath): + def test_delete_archive_file(self, mnt_filepath: GRLDatasets): path1 = Path(pjoin(mnt_filepath.data_src, f"{uuid4().hex}.zip")) # Confirm it doesn't exist, and that delete_archive() doesn't throw @@ -155,7 +157,7 @@ class TestCollectionItemBaseMethods: CollectionItemBase.delete_archive(generic_path=path1) assert not path1.exists() - def test_delete_archive_dir(self, mnt_filepath): + def test_delete_archive_dir(self, mnt_filepath: GRLDatasets): path1 = Path(pjoin(mnt_filepath.data_src, f"{uuid4().hex}")) # Confirm it doesn't exist, and that delete_archive() doesn't throw @@ -174,20 +176,20 @@ class TestCollectionItemBaseMethods: def test_should_archive(self): instance = CollectionItemBase() - with pytest.raises(expected_exception=AttributeError) as cm: - res = instance.should_archive() + with pytest.raises(expected_exception=AttributeError): + _ = instance.should_archive() def test_set_empty(self): instance = CollectionItemBase() - with pytest.raises(expected_exception=AttributeError) as cm: - res = instance.set_empty() + with pytest.raises(expected_exception=AttributeError): + _ = instance.set_empty() def test_valid_archive(self): instance = CollectionItemBase() - with pytest.raises(expected_exception=AttributeError) as cm: - res = instance.valid_archive(generic_path=None, sample=None) + with pytest.raises(expected_exception=AttributeError): + _ = instance.valid_archive(generic_path=None, sample=None) class TestCollectionItemBaseMethodsORM: @@ -197,11 +199,11 @@ class TestCollectionItemBaseMethodsORM: pass @pytest.mark.parametrize("is_partial", [True, False]) - def test_to_archive(self, is_partial): + def test_to_archive(self, is_partial: bool): instance = CollectionItemBase() with pytest.raises(expected_exception=NotImplementedError) as cm: - res = instance.to_archive( + _ = instance.to_archive( ddf=dd.from_pandas(data=pd.DataFrame()), is_partial=is_partial ) assert "Must override" in str(cm.value) diff --git a/tests/incite/test_interval_idx.py b/tests/incite/test_interval_idx.py index 3034c21..03d29ea 100644 --- a/tests/incite/test_interval_idx.py +++ b/tests/incite/test_interval_idx.py @@ -1,4 +1,4 @@ -from datetime import datetime +from datetime import UTC, datetime import pandas as pd @@ -6,8 +6,8 @@ import pandas as pd class TestIntervalIndex: def test_init(self): - start = datetime(year=2000, month=1, day=1) - end = datetime(year=2000, month=1, day=10) + start = datetime(year=2000, month=1, day=1, tzinfo=UTC) + end = datetime(year=2000, month=1, day=10, tzinfo=UTC) iv_r: pd.IntervalIndex = pd.interval_range( start=start, end=end, freq="1d", closed="left" diff --git a/tests/managers/gr/test_business.py b/tests/managers/gr/test_business.py index 3490403..ed141b1 100644 --- a/tests/managers/gr/test_business.py +++ b/tests/managers/gr/test_business.py @@ -2,17 +2,36 @@ from uuid import uuid4 import pytest +from generalresearch.managers.gr.business import ( + BusinessAddressManager, + BusinessBankAccountManager, + BusinessManager, +) +from generalresearch.managers.gr.team import MembershipManager, TeamManager +from generalresearch.models.gr.authentication import GRUser +from generalresearch.models.gr.business import ( + Business, + BusinessAddress, + BusinessBankAccount, + TransferMethod, +) +from generalresearch.pg_helper import PostgresConfig + class TestBusinessBankAccountManager: - def test_init(self, business_bank_account_manager, gr_db): + def test_init( + self, + business_bank_account_manager: BusinessBankAccountManager, + gr_db: PostgresConfig, + ): assert business_bank_account_manager.pg_config == gr_db - def test_create(self, business: Business, business_bank_account_manager): - from generalresearch.models.gr.business import ( - BusinessBankAccount, - TransferMethod, - ) + def test_create( + self, + business: Business, + business_bank_account_manager: BusinessBankAccountManager, + ): instance = business_bank_account_manager.create( business_id=business.id, @@ -33,8 +52,9 @@ class TestBusinessBankAccountManager: class TestBusinessAddressManager: - def test_create(self, business: Business, business_address_manager): - from generalresearch.models.gr.business import BusinessAddress + def test_create( + self, business: Business, business_address_manager: BusinessAddressManager + ): res = business_address_manager.create(uuid=uuid4().hex, business_id=business.id) assert isinstance(res, BusinessAddress) @@ -43,14 +63,13 @@ class TestBusinessAddressManager: class TestBusinessManager: - def test_create(self, business_manager): - from generalresearch.models.gr.business import Business + def test_create(self, business_manager: BusinessManager): instance = business_manager.create_dummy() assert isinstance(instance, Business) assert isinstance(instance.id, int) - def test_get_or_create(self, business_manager): + def test_get_or_create(self, business_manager: BusinessManager): uuid_key = uuid4().hex assert business_manager.get_by_uuid(business_uuid=uuid_key) is None @@ -61,9 +80,10 @@ class TestBusinessManager: ) res = 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): + def test_get_all(self, business_manager: BusinessManager): res1 = business_manager.get_all() assert isinstance(res1, list) @@ -76,7 +96,11 @@ class TestBusinessManager: pass def test_get_by_user_id( - self, business_manager, gr_user, team_manager, membership_manager + self, + business_manager: BusinessManager, + gr_user: GRUser, + team_manager: TeamManager, + membership_manager: MembershipManager, ): res = business_manager.get_by_user_id(user_id=gr_user.id) assert len(res) == 0 @@ -93,7 +117,7 @@ class TestBusinessManager: # Create a Membership for the gr_user to the Team... but it doesn't # matter because the Team doesn't have any Business yet - m1 = membership_manager.create(team=t1, gr_user=gr_user) + _ = membership_manager.create(team=t1, gr_user=gr_user) res = business_manager.get_by_user_id(user_id=gr_user.id) assert len(res) == 0 @@ -113,15 +137,17 @@ class TestBusinessManager: def test_get_uuids_by_user_id(self): pass - def test_get_by_uuid(self, business: Business, business_manager): + def test_get_by_uuid(self, business: Business, business_manager: BusinessManager): instance = business_manager.get_by_uuid(business_uuid=business.uuid) + assert isinstance(instance, Business) assert business.id == instance.id - def test_get_by_id(self, business: Business, business_manager): + def test_get_by_id(self, business: Business, business_manager: BusinessManager): instance = business_manager.get_by_id(business_id=business.id) + assert isinstance(instance, Business) assert business.uuid == instance.uuid - def test_cache_key(self, business): + def test_cache_key(self, business: Business): assert "business:" in business.cache_key # def test_create_raise_on_duplicate(self): diff --git a/tests/managers/gr/test_team.py b/tests/managers/gr/test_team.py index 5e5c565..ae3e1bb 100644 --- a/tests/managers/gr/test_team.py +++ b/tests/managers/gr/test_team.py @@ -1,18 +1,29 @@ +from __future__ import annotations + +from collections.abc import Callable from uuid import uuid4 +from generalresearch.managers.gr.authentication import GRUserManager +from generalresearch.managers.gr.team import MembershipManager, TeamManager +from generalresearch.models.gr.authentication import GRUser +from generalresearch.models.gr.team import Membership, Team +from generalresearch.models.thl.product import Product +from generalresearch.pg_helper import PostgresConfig +from generalresearch.redis_helper import RedisConfig + class TestMembershipManager: - def test_init(self, membership_manager, gr_db): + def test_init(self, membership_manager: MembershipManager, gr_db: PostgresConfig): assert membership_manager.pg_config == gr_db class TestTeamManager: - def test_init(self, team_manager, gr_db): + def test_init(self, team_manager: TeamManager, gr_db: PostgresConfig): assert team_manager.pg_config == gr_db - def test_get_or_create(self, team_manager): + def test_get_or_create(self, team_manager: TeamManager): from generalresearch.models.gr.team import Team new_uuid = uuid4().hex @@ -24,7 +35,7 @@ class TestTeamManager: assert team.uuid == new_uuid assert team.name == "< Unknown >" - def test_get_all(self, team_manager): + def test_get_all(self, team_manager: TeamManager): res1 = team_manager.get_all() assert isinstance(res1, list) @@ -32,16 +43,20 @@ class TestTeamManager: res2 = team_manager.get_all() assert len(res1) == len(res2) - 1 - def test_create(self, team_manager): - from generalresearch.models.gr.team import Team + def test_create(self, team_manager: TeamManager): team: Team = team_manager.create_dummy() assert isinstance(team, Team) assert isinstance(team.id, int) - def test_add_user(self, team, team_manager, gr_um, gr_db, gr_redis_config): - from generalresearch.models.gr.authentication import GRUser - from generalresearch.models.gr.team import Membership + def test_add_user( + self, + team: Team, + team_manager: TeamManager, + gr_um: GRUserManager, + gr_db: PostgresConfig, + gr_redis_config: RedisConfig, + ): user: GRUser = gr_um.create_dummy() @@ -54,25 +69,23 @@ class TestTeamManager: assert len(team.gr_users) assert team.gr_users == [user] - def test_get_by_uuid(self, team_manager): - from generalresearch.models.gr.team import Team + def test_get_by_uuid(self, team_manager: TeamManager): team: Team = team_manager.create_dummy() instance = team_manager.get_by_uuid(team_uuid=team.uuid) assert team.id == instance.id - def test_get_by_id(self, team_manager): - from generalresearch.models.gr.team import Team + def test_get_by_id(self, team_manager: TeamManager): team: Team = team_manager.create_dummy() instance = team_manager.get_by_id(team_id=team.id) assert team.uuid == instance.uuid - def test_get_by_user(self, team, team_manager, gr_um): - from generalresearch.models.gr.authentication import GRUser - from generalresearch.models.gr.team import Team + def test_get_by_user( + self, team: Team, team_manager: TeamManager, gr_um: GRUserManager + ): user: GRUser = gr_um.create_dummy() team_manager.add_user(team=team, gr_user=user) @@ -86,15 +99,12 @@ class TestTeamManager: def test_get_by_user_duplicates( self, - gr_user_token, - gr_user, - membership, + gr_user: GRUser, product_factory: Callable[..., Product], - membership_factory, - team, - thl_web_rr: PostgresConfig, - gr_redis_config, - gr_db, + membership_factory: Callable[..., Membership], + team: Team, + gr_redis_config: RedisConfig, + gr_db: PostgresConfig, ): product_factory(team=team) membership_factory(team=team, gr_user=gr_user) diff --git a/tests/managers/network/test_label.py b/tests/managers/network/test_label.py index bfc7518..71efa95 100644 --- a/tests/managers/network/test_label.py +++ b/tests/managers/network/test_label.py @@ -27,7 +27,7 @@ def ip_label(utc_now) -> IPLabel: provider="GeoNodE", created_at=utc_now, ip=ip, - metadata=IPLabelMetadata(services=["RDP"]) + metadata=IPLabelMetadata(services=["RDP"]), ) @@ -181,7 +181,7 @@ def test_label_cidr_and_ipinfo( ip = fake.ipv6() ip_information_factory(ip=ip, geoname=ip_geoname) # We normalize for storage into ipinfo table - ip_norm, prefix = normalize_ip(ip) + ip_norm, _ = normalize_ip(ip) # Test with a larger network ip_48 = ipaddress.IPv6Network((ip, 48), strict=False) diff --git a/tests/managers/thl/test_contest/test_leaderboard.py b/tests/managers/thl/test_contest/test_leaderboard.py index 07d8d74..3a63075 100644 --- a/tests/managers/thl/test_contest/test_leaderboard.py +++ b/tests/managers/thl/test_contest/test_leaderboard.py @@ -5,7 +5,6 @@ from zoneinfo import ZoneInfo from generalresearch.currency import USDCent from generalresearch.managers.thl.contest_manager import ContestManager -from generalresearch.managers.thl.ledger_manager.ledger import LedgerManager from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager from generalresearch.managers.thl.user_manager.user_manager import UserManager from generalresearch.models.thl.contest.definitions import ( @@ -116,7 +115,9 @@ class TestLeaderboardContestCRUD: assert decision assert reason == ContestEndReason.ENDS_AT - contest_manager.end_contest_if_over(contest=contest, ledger_manager=thl_lm) + contest_manager.end_contest_if_over( + contest=contest, ledger_manager=thl_ledger_manager + ) c: LeaderboardContest = contest_manager.get(contest_uuid=contest.uuid) assert c.status == ContestStatus.COMPLETED diff --git a/tests/managers/thl/test_contest/test_milestone.py b/tests/managers/thl/test_contest/test_milestone.py index a2d575b..e3889bc 100644 --- a/tests/managers/thl/test_contest/test_milestone.py +++ b/tests/managers/thl/test_contest/test_milestone.py @@ -292,7 +292,10 @@ class TestMilestoneContestUserViews: assert len(cs) == 1 contest_manager.enter_milestone_contest( - contest_uuid=c.uuid, user=user, country_iso="us", ledger_manager=thl_lm + contest_uuid=c.uuid, + user=user, + country_iso="us", + ledger_manager=thl_ledger_manager, ) # User isn't eligible anymore diff --git a/tests/managers/thl/test_contest/test_raffle.py b/tests/managers/thl/test_contest/test_raffle.py index b435576..06d4676 100644 --- a/tests/managers/thl/test_contest/test_raffle.py +++ b/tests/managers/thl/test_contest/test_raffle.py @@ -14,6 +14,7 @@ from generalresearch.managers.thl.ledger_manager.exceptions import ( ) from generalresearch.managers.thl.ledger_manager.thl_ledger import ThlLedgerManager from generalresearch.models.thl.contest import ( + Contest, ContestEndCondition, ContestEntryRule, ContestPrize, @@ -40,8 +41,6 @@ class TestRaffleContest: def test_should_end( self, contest: RaffleContest, - thl_ledger_manager: ThlLedgerManager, - contest_manager: ContestManager, ): # contest is active and has no entries should, msg = contest.should_end() @@ -67,7 +66,6 @@ class TestRaffleContestCRUD: self, contest_create: RaffleContestCreate, product_user_wallet_yes: Product, - thl_ledger_manager: ThlLedgerManager, contest_manager: ContestManager, ): c = contest_manager.create( @@ -329,7 +327,7 @@ class TestRaffleContestCRUD: contest_uuid=c.uuid, entry=entry, country_iso="us", - ledger_manager=thl_lm, + ledger_manager=thl_ledger_manager, ) assert "Entry would exceed max amount per user." in str(e.value) @@ -342,7 +340,7 @@ class TestRaffleContestCRUD: contest_uuid=c.uuid, entry=entry, country_iso="us", - ledger_manager=thl_lm, + ledger_manager=thl_ledger_manager, ) assert "Entry would exceed max amount per user per day." in str(e.value) @@ -354,7 +352,7 @@ class TestRaffleContestCRUD: contest_uuid=c.uuid, entry=entry, country_iso="us", - ledger_manager=thl_lm, + ledger_manager=thl_ledger_manager, ) # Then can't anymore @@ -366,7 +364,7 @@ class TestRaffleContestCRUD: contest_uuid=c.uuid, entry=entry, country_iso="us", - ledger_manager=thl_lm, + ledger_manager=thl_ledger_manager, ) assert "Entry would exceed max amount per user per day." in str(e.value) diff --git a/tests/managers/thl/test_ledger/test_lm_accounts.py b/tests/managers/thl/test_ledger/test_lm_accounts.py index 540bea8..7b65b2d 100644 --- a/tests/managers/thl/test_ledger/test_lm_accounts.py +++ b/tests/managers/thl/test_ledger/test_lm_accounts.py @@ -63,7 +63,7 @@ class TestLedgerAccountManagerNoResults: acct_id: UUIDStr, lm: LedgerManager, ): - qn = ":".join([currency, kind, acct_id]) + qn = f"{currency}:{kind}:{acct_id}" # (1) .get_many_ assert lm.get_account_many_(qualified_names=[qn], raise_on_error=False) == [] diff --git a/tests/managers/thl/test_ledger/test_lm_tx.py b/tests/managers/thl/test_ledger/test_lm_tx.py index 13495a7..ce609d6 100644 --- a/tests/managers/thl/test_ledger/test_lm_tx.py +++ b/tests/managers/thl/test_ledger/test_lm_tx.py @@ -24,8 +24,6 @@ class TestLedgerManagerCreateTx: """Confirm that the Permission values that are set on the Ledger Manger allow the Creation action to occur. """ - acct_uuid = uuid4().hex - # (1) With no Permissions defined test_lm = LedgerManager( pg_config=ledger_manager.pg_config, @@ -44,8 +42,6 @@ class TestLedgerManagerCreateTx: def test_create_assertions( self, - ledger_account_debit: LedgerAccount, - ledger_account_credit: LedgerAccount, ledger_manager: LedgerManager, ): with pytest.raises(expected_exception=ValueError) as excinfo: diff --git a/tests/managers/thl/test_ledger/test_thl_lm_accounts.py b/tests/managers/thl/test_ledger/test_thl_lm_accounts.py index dce9116..60eb71c 100644 --- a/tests/managers/thl/test_ledger/test_thl_lm_accounts.py +++ b/tests/managers/thl/test_ledger/test_thl_lm_accounts.py @@ -229,6 +229,7 @@ class TestThlLedgerManagerAccounts: # (1) known account and confirm it comes back res = ledger_manager.get_account(qualified_name=account1.qualified_name) + assert isinstance(res, LedgerAccount) assert account1.model_dump_json() == res.model_dump_json() # (2) known accounts and confirm they both come back @@ -291,6 +292,7 @@ class TestThlLedgerManagerAccounts: assert len(res) == 2 # Confirm an empty array comes back for all unknown qualified names + assert isinstance(ledger_manager.currency, LedgerCurrency) res = ledger_manager.get_accounts_if_exists( qualified_names=[ f"{ledger_manager.currency.value}:bp_wall:{uuid4().hex}" @@ -328,7 +330,7 @@ class TestThlLedgerManagerAccounts: product_uuids=product_uuids ) assert len(res) == len(product_uuids) - assert all([isinstance(i, LedgerAccount) for i in res]) + assert all(isinstance(i, LedgerAccount) for i in res) class TestLedgerAccountManager: @@ -351,10 +353,10 @@ class TestLedgerAccountManager: # First we want to validate that using the get_account method raises # an error for a random LedgerAccount which we know does not exist. with pytest.raises(LedgerAccountDoesntExistError): - lam.get_account(qualified_name=account.qualified_name) + ledger_account_manager.get_account(qualified_name=account.qualified_name) # Now that we know it doesn't exist, get_or_create for it - instance = lam.get_account_or_create(account=account) + instance = ledger_account_manager.get_account_or_create(account=account) # It should always return assert isinstance(instance, LedgerAccount) @@ -364,10 +366,11 @@ class TestLedgerAccountManager: self, user: User, thl_ledger_manager: ThlLedgerManager, - ledger_manager: LedgerManager, ledger_account_manager: LedgerAccountManager, ): + assert isinstance(user.product, Product) + with pytest.raises(LedgerAccountDoesntExistError): ledger_account_manager.get_account( qualified_name=f"test:bp_wallet:{user.product.id}" diff --git a/tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py b/tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py index e4a25a3..b518453 100644 --- a/tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py +++ b/tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py @@ -133,16 +133,17 @@ class TestThlLedgerManagerBPPayout: ) payoutevent_uuid = uuid4().hex - with caplog.at_level(logging.INFO): - with pytest.raises(LedgerTransactionConditionFailedError): - thl_ledger_manager.create_tx_bp_payout( - user.product, - amount=USDCent(10_000), - created=now + timedelta(minutes=2), - skip_one_per_day_check=True, - skip_wallet_balance_check=False, - payoutevent_uuid=payoutevent_uuid, - ) + with caplog.at_level(logging.INFO), pytest.raises( + LedgerTransactionConditionFailedError + ): + thl_ledger_manager.create_tx_bp_payout( + user.product, + amount=USDCent(10_000), + created=now + timedelta(minutes=2), + skip_one_per_day_check=True, + skip_wallet_balance_check=False, + payoutevent_uuid=payoutevent_uuid, + ) assert "failed condition check balance:" in caplog.text thl_ledger_manager.create_tx_bp_payout( @@ -197,17 +198,18 @@ class TestThlLedgerManagerBPPayout: assert balance == int(rand_amount) * -1 # Test some basic assertions - with caplog.at_level(logging.INFO): - with pytest.raises(expected_exception=Exception): - thl_ledger_manager.create_tx_bp_payout( - product=product, - amount=rand_amount, - payoutevent_uuid=uuid4().hex, - created=datetime.now(tz=UTC), - skip_wallet_balance_check=False, - skip_one_per_day_check=False, - skip_flag_check=False, - ) + with caplog.at_level(logging.INFO), pytest.raises( + expected_exception=ValueError + ): + thl_ledger_manager.create_tx_bp_payout( + product=product, + amount=rand_amount, + payoutevent_uuid=uuid4().hex, + created=datetime.now(tz=UTC), + skip_wallet_balance_check=False, + skip_one_per_day_check=False, + skip_flag_check=False, + ) assert "failed condition check >1 tx per day" in caplog.text def test_create_tx_redis_failure( @@ -291,7 +293,7 @@ class TestThlLedgerManagerBPPayout: # Will fail due to multiple per day payoutevent_uuid2 = uuid4().hex with pytest.raises(expected_exception=Exception) as e: - tx = thl_ledger_manager.create_tx_bp_payout( + thl_ledger_manager.create_tx_bp_payout( product=product, amount=rand_amount, payoutevent_uuid=payoutevent_uuid2, @@ -348,7 +350,7 @@ class TestThlLedgerManagerBPPayout: # Create TX will fail on lock exit, after the tx was created! with pytest.raises(expected_exception=Exception) as e: - tx = thl_ledger_manager.create_tx_bp_payout( + thl_ledger_manager.create_tx_bp_payout( product=product, amount=rand_amount, payoutevent_uuid=payoutevent_uuid, @@ -384,7 +386,9 @@ class TestPayoutEventManagerBPPayout: product, rand_amount, now, direction=Direction.CREDIT ) assert thl_ledger_manager.get_account_balance(bp_wallet_account) == rand_amount - brokerage_product_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm) + brokerage_product_payout_event_manager.set_account_lookup_table( + thl_lm=thl_ledger_manager + ) pe = brokerage_product_payout_event_manager.create_bp_payout_event( thl_ledger_manager=thl_ledger_manager, @@ -557,7 +561,7 @@ class TestPayoutEventManagerBPPayout: # Will fail on lock exit, after the tx was created! # But it'll see that the tx was created and so everything will be fine Lock.release = broken_release - pe = brokerage_product_payout_event_manager.create_bp_payout_event( + brokerage_product_payout_event_manager.create_bp_payout_event( thl_ledger_manager=thl_ledger_manager, product=product, created=now, 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 89adb0b..1860d6d 100644 --- a/tests/managers/thl/test_ledger/test_thl_lm_tx.py +++ b/tests/managers/thl/test_ledger/test_thl_lm_tx.py @@ -23,7 +23,6 @@ from generalresearch.models.thl.definitions import ( WALL_ALLOWED_STATUS_STATUS_CODE, ) from generalresearch.models.thl.ledger import ( - AccountType, Direction, LedgerAccount, TransactionType, @@ -287,17 +286,18 @@ class TestThlLedgerTxManager: assert balance == int(rand_amount) * -1 # Test some basic assertions - with caplog.at_level(logging.INFO): - with pytest.raises(expected_exception=Exception): - thl_ledger_manager.create_tx_bp_payout( - product=product, - amount=rand_amount, - payoutevent_uuid=uuid4().hex, - created=datetime.now(tz=UTC), - skip_wallet_balance_check=False, - skip_one_per_day_check=False, - skip_flag_check=False, - ) + with caplog.at_level(logging.INFO), pytest.raises( + expected_exception=ValueError + ): + thl_ledger_manager.create_tx_bp_payout( + product=product, + amount=rand_amount, + payoutevent_uuid=uuid4().hex, + created=datetime.now(tz=UTC), + skip_wallet_balance_check=False, + skip_one_per_day_check=False, + skip_flag_check=False, + ) assert "failed condition check >1 tx per day" in caplog.text def test_create_tx_bp_payout_( @@ -1794,7 +1794,7 @@ class TestThlLedgerManagerAdj: ) thl_ledger_manager.create_tx_bp_payment(session, created=wall1.started) - revenue = ththl_ledger_managerl_lm.get_account_task_complete_revenue() + revenue = thl_ledger_manager.get_account_task_complete_revenue() bp_wallet_account = thl_ledger_manager.get_account_or_create_bp_wallet( user.product ) diff --git a/tests/managers/thl/test_payout.py b/tests/managers/thl/test_payout.py index e6c597b..0f3f103 100644 --- a/tests/managers/thl/test_payout.py +++ b/tests/managers/thl/test_payout.py @@ -727,8 +727,6 @@ class TestBusinessPayoutEventManager: # {"uuid": bp_pe.uuid, "status": PayoutStatus.FAILED}, # ) - assert 1 == 0 - def test_ach_payment( self, mnt_filepath: GRLDatasets, diff --git a/tests/managers/thl/test_survey.py b/tests/managers/thl/test_survey.py index 2c2bf9d..c3ab162 100644 --- a/tests/managers/thl/test_survey.py +++ b/tests/managers/thl/test_survey.py @@ -11,9 +11,6 @@ from generalresearch.managers.thl.buyer import BuyerManager from generalresearch.managers.thl.profiling.question import ( QuestionManager, ) -from generalresearch.managers.thl.profiling.schema import ( - UpkSchemaManager, -) from generalresearch.managers.thl.profiling.uqa import UQAManager from generalresearch.managers.thl.survey import SurveyManager, SurveyStatManager from generalresearch.models import Source @@ -183,7 +180,8 @@ class TestSurvey: ] uqad = {} for uqa in uqas: - for k, _ in uqa.calc_answers.items(): + assert uqa.calc_answers + for k in uqa.calc_answers: if k in qualifying_questions: uqad[k] = uqa uqad[uqa.property_code] = uqa diff --git a/tests/managers/thl/test_user_manager/test_base.py b/tests/managers/thl/test_user_manager/test_base.py index 5822207..8cd83ad 100644 --- a/tests/managers/thl/test_user_manager/test_base.py +++ b/tests/managers/thl/test_user_manager/test_base.py @@ -21,7 +21,7 @@ from generalresearch.managers.thl.user_manager.user_manager import ( UserManager, ) from generalresearch.managers.thl.userhealth import AuditLogManager -from generalresearch.models.thl.product import Product, UserCreateConfig, product +from generalresearch.models.thl.product import Product, UserCreateConfig from generalresearch.models.thl.user import User from generalresearch.pg_helper import PostgresConfig diff --git a/tests/models/network/test_nmap.py b/tests/models/network/test_nmap.py index 5e9f4d0..db39997 100644 --- a/tests/models/network/test_nmap.py +++ b/tests/models/network/test_nmap.py @@ -8,7 +8,7 @@ from generalresearch.managers.network.tool_run import ToolRunManager from generalresearch.models.network.definitions import IPProtocol from generalresearch.models.network.nmap.execute import execute_nmap from generalresearch.models.network.nmap.result import NmapResult, PortState -from generalresearch.models.network.tool_run import NmapRun, Status, ToolClass, ToolName +from generalresearch.models.network.tool_run import NmapRun, ToolClass, ToolName fake = faker.Faker() diff --git a/tests/models/spectrum/test_survey.py b/tests/models/spectrum/test_survey.py index 7ddd407..f97860f 100644 --- a/tests/models/spectrum/test_survey.py +++ b/tests/models/spectrum/test_survey.py @@ -405,3 +405,47 @@ class TestSpectrumSurvey: assert (None, {"c", "d"}) == s.determine_eligibility_soft( {"a": True, "b": True, "c": None, "d": None} ) + + +def test_spectrum_something(spectrum_api_surveys_json: list[str]): + # make sure hashes for 111111 are in db + c1 = SpectrumCondition( + question_id="1001", + value_type=ConditionValueType.LIST, + values=["a", "b", "c"], + negate=False, + logical_operator=LogicalOperator.OR, + ) + c2 = SpectrumCondition( + question_id="1001", + value_type=ConditionValueType.LIST, + values=["a"], + negate=False, + logical_operator=LogicalOperator.OR, + ) + c3 = SpectrumCondition( + question_id="1002", + value_type=ConditionValueType.RANGE, + values=["18-24", "30-32"], + negate=False, + logical_operator=LogicalOperator.OR, + ) + c4 = SpectrumCondition( + question_id="212", + value_type=ConditionValueType.LIST, + values=["23", "24"], + negate=False, + logical_operator=LogicalOperator.OR, + ) + c5 = SpectrumCondition( + question_id="1031", + value_type=ConditionValueType.LIST, + values=["113", "114", "121"], + negate=False, + logical_operator=LogicalOperator.OR, + ) + _conditions = [c1, c2, c3, c4, c5] + + survey = SpectrumSurvey.model_validate_json(spectrum_api_surveys_json[0]) + assert c1.criterion_hash in survey.qualifications + assert c3.criterion_hash in survey.qualifications diff --git a/tests/models/thl/test_product.py b/tests/models/thl/test_product.py index adf276d..880799a 100644 --- a/tests/models/thl/test_product.py +++ b/tests/models/thl/test_product.py @@ -1013,7 +1013,7 @@ class TestProductCache: assert res is None with pytest.raises(expected_exception=AssertionError): product.set_cache( - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, ds=mnt_filepath, client=client_no_amm, bp_pem=brokerage_product_payout_event_manager, @@ -1035,7 +1035,7 @@ class TestProductCache: # Now try again with everything in place product.set_cache( - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, ds=mnt_filepath, client=client_no_amm, bp_pem=brokerage_product_payout_event_manager, -- cgit v1.2.3 From 6469e7e55a53cfe18bd015b3c455ecbbb550cbb9 Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Tue, 1 Sep 2026 12:29:10 -0700 Subject: WIP Business tests, fixture cleanup(s) --- generalresearch/incite/base.py | 4 +- generalresearch/incite/defaults.py | 10 +- generalresearch/managers/__init__.py | 16 -- generalresearch/managers/gr/business.py | 3 +- generalresearch/managers/pollfish/user_pid.py | 2 +- generalresearch/managers/thl/cashout_method.py | 9 +- generalresearch/models/__init__.py | 114 ------------- generalresearch/models/gr/business.py | 13 +- generalresearch/models/gr/definitions.py | 13 ++ generalresearch/models/thl/__init__.py | 18 +-- generalresearch/models/thl/session.py | 8 +- generalresearch/models/thl/task_status.py | 2 +- generalresearch/models/thl/utils.py | 11 ++ generalresearch/models/thl/wallet/__init__.py | 87 ---------- test_utils/conftest.py | 2 +- test_utils/incite/collections/conftest.py | 2 +- test_utils/incite/conftest.py | 10 +- test_utils/incite/mergers/conftest.py | 16 +- test_utils/managers/gr/conftest.py | 28 ---- test_utils/managers/thl/conftest.py | 37 ++++- test_utils/models/conftest.py | 4 +- test_utils/models/contest/conftest.py | 12 +- test_utils/models/gr/conftest.py | 2 +- test_utils/models/ledger/conftest.py | 108 +++++++------ .../incite/collections/test_df_collection_base.py | 6 +- .../collections/test_df_collection_item_base.py | 6 +- tests/incite/test_interval_idx.py | 2 +- tests/managers/gr/test_business.py | 32 ++-- tests/managers/thl/test_ledger/test_lm_accounts.py | 96 ++++++----- tests/managers/thl/test_ledger/test_thl_lm_tx.py | 5 +- tests/managers/thl/test_payout.py | 176 ++++++++++----------- tests/managers/thl/test_session_manager.py | 10 +- tests/models/gr/test_authentication.py | 55 +++---- tests/models/gr/test_business.py | 86 +++++----- tests/models/gr/test_team.py | 6 +- tests/models/test_finance.py | 14 +- tests/models/thl/test_payout.py | 2 +- tests/models/thl/test_product.py | 99 ++++++++---- 38 files changed, 484 insertions(+), 642 deletions(-) create mode 100644 generalresearch/models/gr/definitions.py create mode 100644 generalresearch/models/thl/utils.py (limited to 'tests/incite/test_interval_idx.py') diff --git a/generalresearch/incite/base.py b/generalresearch/incite/base.py index 473a124..a06aac9 100644 --- a/generalresearch/incite/base.py +++ b/generalresearch/incite/base.py @@ -95,7 +95,7 @@ class GRLDatasets(BaseModel): from generalresearch.incite.collections.thl_marketplaces import ( DFCollectionType, ) - from generalresearch.incite.mergers import MergeType + from generalresearch.incite.mergers.base import MergeType assert self.data_src, "data src must be defined" @@ -128,7 +128,7 @@ class GRLDatasets(BaseModel): type.. """ - from generalresearch.incite.mergers import MergeType + from generalresearch.incite.mergers.base import MergeType folder = "mergers" if isinstance(enum_type, MergeType) else "raw/df-collections" assert self.incite is not None diff --git a/generalresearch/incite/defaults.py b/generalresearch/incite/defaults.py index 368b74a..5ee305b 100644 --- a/generalresearch/incite/defaults.py +++ b/generalresearch/incite/defaults.py @@ -3,7 +3,7 @@ from __future__ import annotations from datetime import UTC, datetime from generalresearch.incite.base import GRLDatasets -from generalresearch.incite.collections import DFCollectionType +from generalresearch.incite.collections.base import DFCollectionType from generalresearch.incite.collections.thl_marketplaces import ( InnovateSurveyHistoryCollection, MorningSurveyTimeseriesCollection, @@ -82,7 +82,7 @@ def ledger_df_collection( ds: GRLDatasets, pg_config: PostgresConfig ) -> LedgerDFCollection: return LedgerDFCollection( - offset="12d", + offset="12D", pg_config=pg_config, # thl_web:ledger_transaction - 1st record is 2018-03-14 20:22:17.408232 start=datetime(year=2018, month=3, day=14, hour=0, tzinfo=UTC), @@ -153,7 +153,7 @@ def user_id_product(ds: GRLDatasets) -> UserIdProductMerge: def enriched_session(ds: GRLDatasets) -> EnrichedSessionMerge: return EnrichedSessionMerge( start=datetime(year=2023, month=5, day=1, tzinfo=UTC), - offset="14d", + offset="14D", archive_path=ds.archive_path(enum_type=MergeType.ENRICHED_SESSION), ) @@ -162,7 +162,7 @@ def enriched_wall(ds: GRLDatasets) -> EnrichedWallMerge: return EnrichedWallMerge( # start=datetime(year=2022, month=5, day=1, tzinfo=timezone.utc), start=datetime(year=2023, month=7, day=23, tzinfo=UTC), - offset="14d", + offset="14D", archive_path=ds.archive_path(enum_type=MergeType.ENRICHED_WALL), ) @@ -180,7 +180,7 @@ def pop_ledger(ds: GRLDatasets) -> PopLedgerMerge: return PopLedgerMerge( # thl_web:ledger_transaction - 1st record is 2018-03-14 20:22:17.408232 start=datetime(year=2018, month=3, day=14, hour=0, tzinfo=UTC), - offset="30d", + offset="30D", archive_path=ds.archive_path(enum_type=MergeType.POP_LEDGER), ) diff --git a/generalresearch/managers/__init__.py b/generalresearch/managers/__init__.py index bc745fd..e69de29 100644 --- a/generalresearch/managers/__init__.py +++ b/generalresearch/managers/__init__.py @@ -1,16 +0,0 @@ -def parse_order_by(order_by_str: str) -> str: - """ - Converts django-rest-framework ordering str to mysql clause - :param order_by_str: e.g. 'created,-name' - :return: mysql clause e.g. ORDER BY created ASC, name DESC - """ - fields = order_by_str.split(",") - - order_clause = [] - for field in fields: - if field.startswith("-"): - order_clause.append(f"{field[1:]} DESC") - else: - order_clause.append(f"{field} ASC") - - return "ORDER BY " + ", ".join(order_clause) diff --git a/generalresearch/managers/gr/business.py b/generalresearch/managers/gr/business.py index ef26f30..9bf6ef2 100644 --- a/generalresearch/managers/gr/business.py +++ b/generalresearch/managers/gr/business.py @@ -14,14 +14,13 @@ from generalresearch.managers.base import ( from generalresearch.models.gr.business import ( Business, BusinessBankAccount, - BusinessType, ) +from generalresearch.models.gr.definitions import BusinessType, TransferMethod if TYPE_CHECKING: from generalresearch.models.custom_types import UUIDStr from generalresearch.models.gr.business import ( BusinessAddress, - TransferMethod, ) from generalresearch.models.gr.team import Team diff --git a/generalresearch/managers/pollfish/user_pid.py b/generalresearch/managers/pollfish/user_pid.py index 1068405..f3983cf 100644 --- a/generalresearch/managers/pollfish/user_pid.py +++ b/generalresearch/managers/pollfish/user_pid.py @@ -1,5 +1,5 @@ from generalresearch.managers.marketplace.user_pid import UserPidManager -from generalresearch.models import Source +from generalresearch.models.definitions import Source class PollfishUserPidManager(UserPidManager): diff --git a/generalresearch/managers/thl/cashout_method.py b/generalresearch/managers/thl/cashout_method.py index c12c920..ee86bec 100644 --- a/generalresearch/managers/thl/cashout_method.py +++ b/generalresearch/managers/thl/cashout_method.py @@ -9,15 +9,13 @@ from uuid import UUID, uuid4 from pydantic import NonNegativeInt from generalresearch.managers.base import PostgresManager -from generalresearch.models.thl.wallet.cashout_method import ( - CashoutMethod, -) from generalresearch.models.thl.wallet.definitions import PayoutType if TYPE_CHECKING: from generalresearch.models.thl.user import User from generalresearch.models.thl.wallet.cashout_method import ( CashMailCashoutMethodData, + CashoutMethod, PaypalCashoutMethodData, ) @@ -82,6 +80,7 @@ class CashoutMethodManager(PostgresManager): :return: the uuid of the created cashout method """ # todo: validate shipping address? + from generalresearch.models.thl.wallet.cashout_method import CashoutMethod cm = CashoutMethod( name="Cash in Mail", @@ -126,6 +125,8 @@ class CashoutMethodManager(PostgresManager): :param user: :return: the uuid of the created cashout method """ + from generalresearch.models.thl.wallet.cashout_method import CashoutMethod + cm = CashoutMethod( name="PayPal", description="Cashout via PayPal", @@ -290,6 +291,8 @@ class CashoutMethodManager(PostgresManager): # The data column here is inconsistent. Pulling keys from the mysql 'data' col # and putting them into the base level. Renamed so that we don't overwrite # a col called "data" within the "_data_" field. + from generalresearch.models.thl.wallet.cashout_method import CashoutMethod + for k in list(x["_data_"].keys()): if k in CashoutMethod.model_fields: x[k] = x["_data_"].pop(k) diff --git a/generalresearch/models/__init__.py b/generalresearch/models/__init__.py index c0348d7..e69de29 100644 --- a/generalresearch/models/__init__.py +++ b/generalresearch/models/__init__.py @@ -1,114 +0,0 @@ -from __future__ import annotations - -from enum import IntEnum, StrEnum - -from generalresearch.utils.enum import ReprEnumMeta - - -class Source(StrEnum, metaclass=ReprEnumMeta): - # The external marketplace, or the source of the survey / work. - # Max length of the value is 2. - GRS = "g" - CINT = "c" - DALIA = "a" # deprecated - DYNATA = "d" - ETX = "et" - FULL_CIRCLE = "f" - INNOVATE = "i" - LUCID = "l" - MORNING_CONSULT = "m" - OPEN_LABS = "n" - POLLFISH = "o" - PRECISION = "e" - PRODEGE_USER = "r" # deprecated - PRODEGE = "pr" # using 'r' for vendor_wall - PULLEY = "p" # deprecated - REPDATA = "rd" # using 'q' for vendor_wall - SAGO = "h" - SPECTRUM = "s" - TESTING = "t" # Used internally for testing - TESTING2 = "u" # Used internally for testing - WXET = "w" - - -class DebitKey(IntEnum, metaclass=ReprEnumMeta): - # The debit key for marketplaces - CINT = 8 - DALIA = 9 - DYNATA = 6 - # ETX = None - FULL_CIRCLE = 15 - INNOVATE = 7 - LUCID = 0 - MORNING_CONSULT = 12 - # OPEN_LABS = None - POLLFISH = 13 - PRECISION = 14 - PRODEGE = 11 - SAGO = 10 - SPECTRUM = 5 - # WXET = None - - -class DeviceType(IntEnum, metaclass=ReprEnumMeta): - UNKNOWN = 0 - MOBILE = 1 - DESKTOP = 2 - TABLET = 3 - - -class LogicalOperator(StrEnum, metaclass=ReprEnumMeta): - OR = "OR" - AND = "AND" - # There is currently no use case for NOT. See MarketplaceCondition.explain_not - NOT = "NOT" - - -class TaskStatus(StrEnum, metaclass=ReprEnumMeta): - # A survey is live if it is open and, given all conditions are met, is - # possible to send in traffic. All other statuses are just variants of - # NOT Live (not accepting traffic) - LIVE = "LIVE" - - # This is a generic NOT Live status. A marketplace may use other more - # specific statuses but in practice they don't matter because all we care - # about is if the task is LIVE. - NOT_LIVE = "NOT_LIVE" - - # We need a status to mark if a survey we thought was live does not come - # back from the API, we'll mark it as NOT_FOUND. - NOT_FOUND = "NOT_FOUND" - - -class TaskCalculationType(StrEnum): - COMPLETES = "COMPLETES" - STARTS = "STARTS" - - @classmethod - def from_api(cls, v: str) -> TaskCalculationType: - return { - "complete": cls.COMPLETES, - "completes": cls.COMPLETES, - "survey start": cls.STARTS, - "survey starts": cls.STARTS, - "start": cls.STARTS, - "prescreens": cls.STARTS, - "prescreen": cls.STARTS, - }[v.lower()] - - @classmethod - def prodege_from_api(cls, v: int) -> TaskCalculationType: - return {1: cls.COMPLETES, 2: cls.STARTS}[v] - - @classmethod - def innovate_from_api(cls, v: int) -> TaskCalculationType: - return {0: cls.COMPLETES, 1: cls.STARTS}[v] - - -class URLQueryKey(StrEnum, metaclass=ReprEnumMeta): - PRODUCT_ID = "39057c8b" - PRODUCT_USER_ID = "c184efc0" - SESSION_ID = "0bb50182" - - -MAX_INT32 = 2**31 diff --git a/generalresearch/models/gr/business.py b/generalresearch/models/gr/business.py index e11c54d..c6d3468 100644 --- a/generalresearch/models/gr/business.py +++ b/generalresearch/models/gr/business.py @@ -4,7 +4,6 @@ import json import logging import os from datetime import UTC, datetime -from enum import Enum, StrEnum from pathlib import Path from typing import TYPE_CHECKING from uuid import uuid4 @@ -29,11 +28,11 @@ from generalresearch.models.custom_types import ( UUIDStr, UUIDStrCoerce, ) +from generalresearch.models.gr.definitions import BusinessType, TransferMethod from generalresearch.models.gr.team import Team from generalresearch.models.thl.finance import BusinessBalances, POPFinancial from generalresearch.models.thl.ledger import OrderBy from generalresearch.utils.aggregation import group_by_year -from generalresearch.utils.enum import ReprEnumMeta if TYPE_CHECKING: from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge @@ -69,16 +68,6 @@ if TYPE_CHECKING: from generalresearch.models.thl.product import Product -class TransferMethod(Enum, metaclass=ReprEnumMeta): - ACH = 0 - WIRE = 1 - - -class BusinessType(StrEnum, metaclass=ReprEnumMeta): - INDIVIDUAL = "i" - COMPANY = "c" - - class BusinessBankAccount(BaseModel): model_config = ConfigDict( use_enum_values=True, diff --git a/generalresearch/models/gr/definitions.py b/generalresearch/models/gr/definitions.py new file mode 100644 index 0000000..2e06c03 --- /dev/null +++ b/generalresearch/models/gr/definitions.py @@ -0,0 +1,13 @@ +from enum import Enum, StrEnum + +from generalresearch.utils.enum import ReprEnumMeta + + +class TransferMethod(Enum, metaclass=ReprEnumMeta): + ACH = 0 + WIRE = 1 + + +class BusinessType(StrEnum, metaclass=ReprEnumMeta): + INDIVIDUAL = "i" + COMPANY = "c" diff --git a/generalresearch/models/thl/__init__.py b/generalresearch/models/thl/__init__.py index 7f2b8a9..45278f8 100644 --- a/generalresearch/models/thl/__init__.py +++ b/generalresearch/models/thl/__init__.py @@ -1,14 +1,12 @@ -from decimal import Decimal - # from generalresearch.models.thl.finance import ( # POPFinancial, # ProductBalances, # ) # from generalresearch.models.thl.payout import ( -# BrokerageProductPayoutEvent, +# # BrokerageProductPayoutEvent, # PayoutEvent, # ) -from generalresearch.models.thl.product import Product +# from generalresearch.models.thl.product import Product # _ = ( # Product, @@ -18,16 +16,6 @@ from generalresearch.models.thl.product import Product # POPFinancial, # ) -Product.model_rebuild() +# Product.model_rebuild() # PayoutEvent.model_rebuild() # BrokerageProductPayoutEvent.model_rebuild() - - -def decimal_to_int_cents(usd: Decimal | None) -> int | None: - return round(usd * 100) if usd is not None else None - - -def int_cents_to_decimal(value: int | None, decimals: int = 2) -> Decimal | None: - if value is None: - return None - return (Decimal(value) / Decimal(100)).quantize(Decimal(10) ** -decimals) diff --git a/generalresearch/models/thl/session.py b/generalresearch/models/thl/session.py index 404cff7..65b885e 100644 --- a/generalresearch/models/thl/session.py +++ b/generalresearch/models/thl/session.py @@ -19,10 +19,6 @@ from pydantic import ( ) from generalresearch.models.definitions import Source -from generalresearch.models.thl import ( - decimal_to_int_cents, - int_cents_to_decimal, -) from generalresearch.models.thl.definitions import ( WALL_ALLOWED_STATUS_CODE_1_2, WALL_ALLOWED_STATUS_STATUS_CODE, @@ -32,6 +28,10 @@ from generalresearch.models.thl.definitions import ( WallAdjustedStatus, WallStatusCode2, ) +from generalresearch.models.thl.utils import ( + decimal_to_int_cents, + int_cents_to_decimal, +) if TYPE_CHECKING: from generalresearch.managers.thl.ledger_manager.thl_ledger import ( diff --git a/generalresearch/models/thl/task_status.py b/generalresearch/models/thl/task_status.py index 817f4c5..6cff884 100644 --- a/generalresearch/models/thl/task_status.py +++ b/generalresearch/models/thl/task_status.py @@ -13,7 +13,6 @@ from pydantic import ( model_validator, ) -from generalresearch.models.thl import decimal_to_int_cents from generalresearch.models.thl.definitions import ( SessionAdjustedStatus, SessionStatusCode2, @@ -25,6 +24,7 @@ from generalresearch.models.thl.payout_format import ( PayoutFormatOptionalField, ) from generalresearch.models.thl.session import WallOut +from generalresearch.models.thl.utils import decimal_to_int_cents if TYPE_CHECKING: from generalresearch.models.custom_types import ( diff --git a/generalresearch/models/thl/utils.py b/generalresearch/models/thl/utils.py new file mode 100644 index 0000000..3e14065 --- /dev/null +++ b/generalresearch/models/thl/utils.py @@ -0,0 +1,11 @@ +from decimal import Decimal + + +def decimal_to_int_cents(usd: Decimal | None) -> int | None: + return round(usd * 100) if usd is not None else None + + +def int_cents_to_decimal(value: int | None, decimals: int = 2) -> Decimal | None: + if value is None: + return None + return (Decimal(value) / Decimal(100)).quantize(Decimal(10) ** -decimals) diff --git a/generalresearch/models/thl/wallet/__init__.py b/generalresearch/models/thl/wallet/__init__.py index 2d1eb8d..e69de29 100644 --- a/generalresearch/models/thl/wallet/__init__.py +++ b/generalresearch/models/thl/wallet/__init__.py @@ -1,87 +0,0 @@ -from enum import StrEnum - -from generalresearch.utils.enum import ReprEnumMeta - - -class PayoutType(StrEnum, metaclass=ReprEnumMeta): - """ - The method in which the requested payout is delivered. - """ - - # The max size of the db field that holds this value is 14, so please - # don't add new values longer than that! - - # User is paid out to their personal PayPal email address - PAYPAL = "PAYPAL" - # User is paid out via a Tango Gift Card - TANGO = "TANGO" - # DWOLLA - DWOLLA = "DWOLLA" - # A payment is made to a bank account using ACH - ACH = "ACH" - # A payment is made to a bank account using ACH - WIRE = "WIRE" - # A payment is made in cash and mailed to the user. - CASH_IN_MAIL = "CASH_IN_MAIL" - # A payment is made as a prize with some monetary value - PRIZE = "PRIZE" - - # This is used to designate either AMT_BONUS or AMT_HIT - AMT = "AMT" - # Amazon Mechanical Turk as a Bonus - AMT_BONUS = "AMT_BONUS" - # Amazon Mechanical Turk for a HIT - AMT_HIT = "AMT_ASSIGNMENT" - AMT_ASSIGNMENT = "AMT_ASSIGNMENT" - - -class Currency(StrEnum): - # United States Dollar - USD = "USD" - # Canadian Dollar - CAD = "CAD" - # British Pound Sterling - GBP = "GBP" - # Euro - EUR = "EUR" - # Indian Rupee - INR = "INR" - # Australian Dollar - AUD = "AUD" - # Polish Zloty - PLN = "PLN" - # Swedish Krona - SEK = "SEK" - # Singapore Dollar - SGD = "SGD" - # Mexican Peso - MXN = "MXN" - - -CURRENCY_FORMATTER = { - "USD": lambda x: f"${x / 100:,.2f}", - "CAD": lambda x: f"${x / 100:,.2f} CAD", - "GBP": lambda x: f"{x / 100:,.2f} £", - "EUR": lambda x: f"€{x / 100:,.2f}", - "INR": lambda x: f"₹{x / 100:,.2f}", - "AUD": lambda x: f"${x / 100:,.2f} AUD", - "PLN": lambda x: f"{x / 100:,.2f} zł", - "SEK": lambda x: f"{x / 100:,.2f} kr", - "SGD": lambda x: f"${x / 100:,.2f} SGD", - "MXN": lambda x: f"${x / 100:,.2f} MXN", -} - -# The max value user can redeem in one go in foreign currencies. should be < $250 -# in order to avoid exchange rate issues -CURRENCY_MAX_VALUE = { - "USD": 250, - "CAD": 200, - "GBP": 100, - "EUR": 100, - "INR": 10000, - "AUD": 200, - "PLN": 500, - "SEK": 1000, - "SGD": 200, - "MXN": 4000, -} diff --git a/test_utils/conftest.py b/test_utils/conftest.py index 397d98f..daf6b43 100644 --- a/test_utils/conftest.py +++ b/test_utils/conftest.py @@ -342,7 +342,7 @@ def delete_df_collection( thl_web_rw: PostgresConfig, create_main_accounts: Callable[..., None] ) -> Callable[..., None]: - from generalresearch.incite.collections import ( + from generalresearch.incite.collections.base import ( DFCollection, DFCollectionType, ) diff --git a/test_utils/incite/collections/conftest.py b/test_utils/incite/collections/conftest.py index f490e14..499f90b 100644 --- a/test_utils/incite/collections/conftest.py +++ b/test_utils/incite/collections/conftest.py @@ -197,7 +197,7 @@ def df_collection( utc_90days_ago: datetime, thl_web_rr: PostgresConfig, ) -> DFCollection: - from generalresearch.incite.collections import DFCollection + from generalresearch.incite.collections.base import DFCollection start = utc_90days_ago.replace(microsecond=0) diff --git a/test_utils/incite/conftest.py b/test_utils/incite/conftest.py index 2968d18..bcf0511 100644 --- a/test_utils/incite/conftest.py +++ b/test_utils/incite/conftest.py @@ -16,11 +16,11 @@ from faker import Faker if TYPE_CHECKING: from generalresearch.config import GRLBaseSettings from generalresearch.incite.base import GRLDatasets - from generalresearch.incite.collections import ( + from generalresearch.incite.collections.base import ( DFCollectionItem, DFCollectionType, ) - from generalresearch.incite.mergers import MergeType + from generalresearch.incite.mergers.base import MergeType from generalresearch.models.admin.request import ( ReportRequest, ) @@ -131,14 +131,14 @@ def duration() -> timedelta | None: @pytest.fixture def df_collection_data_type() -> DFCollectionType: - from generalresearch.incite.collections import DFCollectionType + from generalresearch.incite.collections.base import DFCollectionType return DFCollectionType.TEST @pytest.fixture def merge_type() -> MergeType: - from generalresearch.incite.mergers import MergeType + from generalresearch.incite.mergers.base import MergeType return MergeType.TEST @@ -156,7 +156,7 @@ def incite_item_factory( observations: int = 3, user: User | None = None, ): - from generalresearch.incite.collections import ( + from generalresearch.incite.collections.base import ( DFCollection, DFCollectionType, ) diff --git a/test_utils/incite/mergers/conftest.py b/test_utils/incite/mergers/conftest.py index 4eb3f2d..fb95c81 100644 --- a/test_utils/incite/mergers/conftest.py +++ b/test_utils/incite/mergers/conftest.py @@ -58,7 +58,7 @@ def pop_ledger_merge( duration: timedelta, ) -> PopLedgerMerge: - from generalresearch.incite.mergers import MergeType + from generalresearch.incite.mergers.base import MergeType from generalresearch.incite.mergers.pop_ledger import PopLedgerMerge return PopLedgerMerge( @@ -88,7 +88,7 @@ def ym_survey_wall_merge( mnt_filepath: GRLDatasets, start: datetime, ) -> YMSurveyWallMerge: - from generalresearch.incite.mergers import MergeType + from generalresearch.incite.mergers.base import MergeType from generalresearch.incite.mergers.ym_survey_wall import YMSurveyWallMerge return YMSurveyWallMerge( @@ -119,7 +119,7 @@ def ym_wall_summary_merge( duration: timedelta, start: datetime, ) -> YMWallSummaryMerge: - from generalresearch.incite.mergers import MergeType + from generalresearch.incite.mergers.base import MergeType from generalresearch.incite.mergers.ym_wall_summary import YMWallSummaryMerge return YMWallSummaryMerge( @@ -155,7 +155,7 @@ def enriched_session_merge( duration: timedelta, start: datetime, ) -> EnrichedSessionMerge: - from generalresearch.incite.mergers import MergeType + from generalresearch.incite.mergers.base import MergeType from generalresearch.incite.mergers.foundations.enriched_session import ( EnrichedSessionMerge, ) @@ -175,7 +175,7 @@ def enriched_task_adjust_merge( duration: timedelta, start: datetime, ) -> EnrichedTaskAdjustMerge: - from generalresearch.incite.mergers import MergeType + from generalresearch.incite.mergers.base import MergeType from generalresearch.incite.mergers.foundations.enriched_task_adjust import ( EnrichedTaskAdjustMerge, ) @@ -197,7 +197,7 @@ def enriched_wall_merge( duration: timedelta, start: datetime, ) -> EnrichedWallMerge: - from generalresearch.incite.mergers import MergeType + from generalresearch.incite.mergers.base import MergeType from generalresearch.incite.mergers.foundations.enriched_wall import ( EnrichedWallMerge, ) @@ -217,7 +217,7 @@ def user_id_product_merge( offset: str, start: datetime, ) -> UserIdProductMerge: - from generalresearch.incite.mergers import MergeType + from generalresearch.incite.mergers.base import MergeType from generalresearch.incite.mergers.foundations.user_id_product import ( UserIdProductMerge, ) @@ -243,7 +243,7 @@ def merge_collection( duration: timedelta, start: datetime, ): - from generalresearch.incite.mergers import MergeCollection + from generalresearch.incite.mergers.base import MergeCollection return MergeCollection( merge_type=merge_type, diff --git a/test_utils/managers/gr/conftest.py b/test_utils/managers/gr/conftest.py index 5392c69..a7fa9e9 100644 --- a/test_utils/managers/gr/conftest.py +++ b/test_utils/managers/gr/conftest.py @@ -9,7 +9,6 @@ import pytest import redis import redis.asyncio as redis_async from pydantic import PostgresDsn -from redis import Redis from generalresearch.managers.gr.business import ( BusinessAddressManager, @@ -30,33 +29,6 @@ def gr_redis_config_db() -> str: return str(randint(99, 1_023)) -@pytest.fixture(scope="session") -def gr_redis(settings: GRLBaseSettings) -> Redis: - assert "unittest" in str(settings.testing_redis) or "127.0.0.1" in str( - settings.testing_redis - ) - return Redis.from_url( - url=str(settings.gr_redis), - decode_responses=True, - socket_timeout=settings.redis_timeout, - socket_connect_timeout=settings.redis_timeout, - ) - - -@pytest.fixture -def gr_redis_async(settings: GRLBaseSettings) -> redis_async.Redis: - assert "unittest" in str(settings.testing_redis) or "127.0.0.1" in str( - settings.testing_redis - ) - - return redis_async.Redis.from_url( - str(settings.testing_redis), - decode_responses=True, - socket_timeout=0.20, - socket_connect_timeout=0.20, - ) - - @pytest.fixture(scope="session") def gr_redis_config( settings: GRLBaseSettings, gr_redis_config_db: str diff --git a/test_utils/managers/thl/conftest.py b/test_utils/managers/thl/conftest.py index af3fd23..391b74c 100644 --- a/test_utils/managers/thl/conftest.py +++ b/test_utils/managers/thl/conftest.py @@ -1,9 +1,12 @@ from __future__ import annotations -from collections.abc import Callable +import subprocess +from collections.abc import Callable, Generator +from random import randint from typing import TYPE_CHECKING import pytest +import redis from pydantic import PostgresDsn from generalresearch.managers.base import Permission @@ -59,14 +62,40 @@ def thl_web_rw(thl_web_rr: PostgresConfig) -> PostgresConfig: @pytest.fixture(scope="session") -def thl_redis_config(settings: GRLBaseSettings) -> RedisConfig: - return RedisConfig( - dsn=settings.thl_redis, +def thl_redis_config_db() -> str: + return str(randint(99, 1_023)) + + +@pytest.fixture(scope="session") +def thl_redis_config( + settings: GRLBaseSettings, thl_redis_config_db: str +) -> Generator[RedisConfig]: + assert "unittest" in str(settings.testing_redis) or "127.0.0.1" in str( + settings.testing_redis + ) + + uri = f"redis://{settings.testing_redis}/{thl_redis_config_db}" + + res = subprocess.run( + ["redis-cli", "-u", uri, "SET", "jenkins_lock", "1", "NX", "EX", "3600"], + check=True, + text=True, + capture_output=True, + ) + + if res.stdout.strip() != "OK": + raise ValueError("Redis already locked... aborting.") + + yield RedisConfig( + dsn=uri, decode_responses=True, socket_timeout=settings.redis_timeout, socket_connect_timeout=settings.redis_timeout, ) + r = redis.from_url(uri) + r.flushdb() + @pytest.fixture(scope="session") def payout_event_manager( diff --git a/test_utils/models/conftest.py b/test_utils/models/conftest.py index 089f2e6..ed4da08 100644 --- a/test_utils/models/conftest.py +++ b/test_utils/models/conftest.py @@ -370,7 +370,7 @@ def product_amt_true( @pytest.fixture def bp_payout_factory( - thl_lm: ThlLedgerManager, + thl_ledger_manager: ThlLedgerManager, product_manager: ProductManager, business_payout_event_manager: BusinessPayoutEventManager, ) -> Callable[..., BrokerageProductPayoutEvent]: @@ -389,7 +389,7 @@ def bp_payout_factory( amount = amount or USDCent(randint(1, 99_99)) return business_payout_event_manager.create_bp_payout_event( - thl_ledger_manager=thl_lm, + thl_ledger_manager=thl_ledger_manager, product=product, amount=amount, ext_ref_id=ext_ref_id or uuid4().hex, diff --git a/test_utils/models/contest/conftest.py b/test_utils/models/contest/conftest.py index 91425dc..18a8e5f 100644 --- a/test_utils/models/contest/conftest.py +++ b/test_utils/models/contest/conftest.py @@ -275,24 +275,26 @@ def user_with_money( request: Request, user_factory: Callable[..., User], product_user_wallet_yes: Product, - thl_lm: ThlLedgerManager, + thl_ledger_manager: 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) + wallet = thl_ledger_manager.get_account_or_create_user_wallet(user) + balance = thl_ledger_manager.get_account_balance(wallet) todo = min_balance - balance if todo > 0: # # Put money in user's wallet - thl_lm.create_tx_user_bonus( + thl_ledger_manager.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)}") + print( + f"wallet balance: {thl_ledger_manager.get_user_wallet_balance(user=user)}" + ) return user diff --git a/test_utils/models/gr/conftest.py b/test_utils/models/gr/conftest.py index 6c1877a..e493f20 100644 --- a/test_utils/models/gr/conftest.py +++ b/test_utils/models/gr/conftest.py @@ -23,8 +23,8 @@ if TYPE_CHECKING: Business, BusinessAddress, BusinessBankAccount, - TransferMethod, ) + from generalresearch.models.gr.definitions import TransferMethod from generalresearch.models.gr.team import Membership, Team from generalresearch.pg_helper import PostgresConfig from generalresearch.redis_helper import RedisConfig diff --git a/test_utils/models/ledger/conftest.py b/test_utils/models/ledger/conftest.py index 8437c7f..31e5eb4 100644 --- a/test_utils/models/ledger/conftest.py +++ b/test_utils/models/ledger/conftest.py @@ -65,7 +65,7 @@ if TYPE_CHECKING: @pytest.fixture def ledger_account( - request: Request, lm: LedgerManager, currency: LedgerCurrency + request: Request, ledger_manager: LedgerManager, currency: LedgerCurrency ) -> LedgerAccount: from generalresearch.models.thl.ledger import ( AccountType, @@ -87,14 +87,14 @@ def ledger_account( account_type=account_type, normal_balance=direction, ) - return lm.create_account(account=acct_model) + return ledger_manager.create_account(account=acct_model) @pytest.fixture def ledger_account_factory( request: Request, - thl_lm: ThlLedgerManager, - lm: LedgerManager, + thl_ledger_manager: ThlLedgerManager, + ledger_manager: LedgerManager, currency: LedgerCurrency, ) -> Callable[..., LedgerAccount]: @@ -109,7 +109,7 @@ def ledger_account_factory( account_type: AccountType = AccountType.CASH, direction: Direction = Direction.CREDIT, ) -> LedgerAccount: - thl_lm.get_account_or_create_bp_wallet(product=product) + thl_ledger_manager.get_account_or_create_bp_wallet(product=product) acct_uuid = uuid4().hex qn = f"{currency}:{account_type}:{acct_uuid}" @@ -121,14 +121,14 @@ def ledger_account_factory( account_type=account_type, normal_balance=direction, ) - return lm.create_account(account=acct_model) + return ledger_manager.create_account(account=acct_model) return _inner @pytest.fixture def ledger_account_credit( - request: Request, lm: LedgerManager, currency: LedgerCurrency + request: Request, ledger_manager: LedgerManager, currency: LedgerCurrency ) -> LedgerAccount: from generalresearch.models.thl.ledger import AccountType, Direction @@ -146,12 +146,12 @@ def ledger_account_credit( account_type=account_type, normal_balance=Direction.CREDIT, ) - return lm.create_account(account=acct_model) + return ledger_manager.create_account(account=acct_model) @pytest.fixture def ledger_account_debit( - request: Request, lm: LedgerManager, currency: LedgerCurrency + request: Request, ledger_manager: LedgerManager, currency: LedgerCurrency ) -> LedgerAccount: from generalresearch.models.thl.ledger import AccountType, Direction @@ -169,11 +169,11 @@ def ledger_account_debit( account_type=account_type, normal_balance=Direction.DEBIT, ) - return lm.create_account(account=acct_model) + return ledger_manager.create_account(account=acct_model) @pytest.fixture -def tag(request: Request, lm: LedgerManager) -> str: +def tag(request: Request) -> str: from generalresearch.currency import LedgerCurrency return ( @@ -194,11 +194,11 @@ def bp_payout_event( product: Product, usd_cent: USDCent, business_payout_event_manager: BusinessPayoutEventManager, - thl_lm: ThlLedgerManager, + thl_ledger_manager: ThlLedgerManager, ) -> BrokerageProductPayoutEvent: return business_payout_event_manager.create_bp_payout_event( - thl_ledger_manager=thl_lm, + thl_ledger_manager=thl_ledger_manager, product=product, amount=usd_cent, skip_wallet_balance_check=True, @@ -209,7 +209,7 @@ def bp_payout_event( @pytest.fixture def bp_payout_event_factory( brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager, - thl_lm: ThlLedgerManager, + thl_ledger_manager: ThlLedgerManager, ) -> Callable[..., BrokerageProductPayoutEvent]: def _inner( @@ -217,7 +217,7 @@ def bp_payout_event_factory( ) -> BrokerageProductPayoutEvent: return brokerage_product_payout_event_manager.create_bp_payout_event( - thl_ledger_manager=thl_lm, + thl_ledger_manager=thl_ledger_manager, product=product, amount=usd_cent, ext_ref_id=ext_ref_id, @@ -229,10 +229,12 @@ def bp_payout_event_factory( @pytest.fixture -def currency(lm: LedgerManager) -> LedgerCurrency: +def currency(ledger_manager: 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 + assert ( + ledger_manager.currency + ), "LedgerManager must have a currency specified for these tests" + return ledger_manager.currency @pytest.fixture @@ -252,7 +254,7 @@ def ledger_tx( tag: str, currency: LedgerCurrency, tx_metadata: dict[str, str] | None, - lm: LedgerManager, + ledger_manager: LedgerManager, ) -> LedgerTransaction: from generalresearch.models.thl.ledger import Direction, LedgerEntry @@ -271,12 +273,12 @@ def ledger_tx( ), ] - return lm.create_tx(entries=entries, tag=tag, metadata=tx_metadata) + return ledger_manager.create_tx(entries=entries, tag=tag, metadata=tx_metadata) @pytest.fixture def create_main_accounts( - lm: LedgerManager, currency: LedgerCurrency + ledger_manager: LedgerManager, currency: LedgerCurrency ) -> Callable[..., None]: def _inner() -> None: @@ -291,9 +293,9 @@ def create_main_accounts( qualified_name=f"{currency.value}:revenue:task_complete", normal_balance=Direction.CREDIT, account_type=AccountType.REVENUE, - currency=lm.currency, + currency=ledger_manager.currency, ) - lm.get_account_or_create(account=account) + ledger_manager.get_account_or_create(account=account) account = LedgerAccount( display_name="Operating Cash Account", @@ -303,7 +305,7 @@ def create_main_accounts( currency=currency, ) - lm.get_account_or_create(account=account) + ledger_manager.get_account_or_create(account=account) return _inner @@ -327,7 +329,7 @@ def delete_ledger_db(thl_web_rw: PostgresManager) -> Callable[..., None]: @pytest.fixture def wipe_main_accounts( - thl_web_rw: PostgresManager, lm: LedgerManager, currency: LedgerCurrency + thl_web_rw: PostgresManager, ledger_manager: LedgerManager, currency: LedgerCurrency ) -> Callable[..., None]: def _inner() -> None: @@ -397,7 +399,9 @@ def wipe_main_accounts( @pytest.fixture -def account_cash(lm: LedgerManager, currency: LedgerCurrency) -> LedgerAccount: +def account_cash( + ledger_manager: LedgerManager, currency: LedgerCurrency +) -> LedgerAccount: from generalresearch.models.thl.ledger import ( AccountType, Direction, @@ -411,12 +415,12 @@ def account_cash(lm: LedgerManager, currency: LedgerCurrency) -> LedgerAccount: account_type=AccountType.CASH, currency=currency, ) - return lm.get_account_or_create(account=account) + return ledger_manager.get_account_or_create(account=account) @pytest.fixture def account_revenue_task_complete( - lm: LedgerManager, currency: LedgerCurrency + ledger_manager: LedgerManager, currency: LedgerCurrency ) -> LedgerAccount: from generalresearch.models.thl.ledger import ( AccountType, @@ -431,11 +435,13 @@ def account_revenue_task_complete( account_type=AccountType.REVENUE, currency=currency, ) - return lm.get_account_or_create(account=account) + return ledger_manager.get_account_or_create(account=account) @pytest.fixture -def account_expense_tango(lm: LedgerManager, currency: LedgerCurrency) -> LedgerAccount: +def account_expense_tango( + ledger_manager: LedgerManager, currency: LedgerCurrency +) -> LedgerAccount: from generalresearch.models.thl.ledger import ( AccountType, Direction, @@ -449,12 +455,12 @@ def account_expense_tango(lm: LedgerManager, currency: LedgerCurrency) -> Ledger account_type=AccountType.EXPENSE, currency=currency, ) - return lm.get_account_or_create(account=account) + return ledger_manager.get_account_or_create(account=account) @pytest.fixture def user_account_user_wallet( - lm: LedgerManager, user: User, currency: LedgerCurrency + ledger_manager: LedgerManager, user: User, currency: LedgerCurrency ) -> LedgerAccount: from generalresearch.models.thl.ledger import ( AccountType, @@ -471,12 +477,12 @@ def user_account_user_wallet( reference_uuid=user.uuid, currency=currency, ) - return lm.get_account_or_create(account=account) + return ledger_manager.get_account_or_create(account=account) @pytest.fixture def product_account_bp_wallet( - lm: LedgerManager, product: Product, currency: LedgerCurrency + ledger_manager: LedgerManager, product: Product, currency: LedgerCurrency ) -> LedgerAccount: from generalresearch.models.thl.ledger import ( AccountType, @@ -495,13 +501,13 @@ def product_account_bp_wallet( "currency": currency, } ) - return lm.get_account_or_create(account=account) + return ledger_manager.get_account_or_create(account=account) @pytest.fixture def setup_accounts( product_factory: Callable[..., Product], - lm: LedgerManager, + ledger_manager: LedgerManager, user: User, currency: LedgerCurrency, ) -> Callable[..., None]: @@ -524,7 +530,7 @@ def setup_accounts( reference_uuid=p1.uuid, currency=currency, ) - lm.get_account_or_create(account=account) + ledger_manager.get_account_or_create(account=account) account = LedgerAccount.model_validate( { @@ -537,7 +543,7 @@ def setup_accounts( "currency": currency, } ) - lm.get_account_or_create(account=account) + ledger_manager.get_account_or_create(account=account) # BP's wallet, user's wallet, and a revenue from their commissions account. p2 = product_factory() @@ -550,7 +556,7 @@ def setup_accounts( reference_uuid=p2.uuid, currency=currency, ) - lm.get_account_or_create(account) + ledger_manager.get_account_or_create(account) account = LedgerAccount( display_name=f"{p2.name} Wallet", @@ -561,7 +567,7 @@ def setup_accounts( reference_uuid=p2.uuid, currency=currency, ) - lm.get_account_or_create(account) + ledger_manager.get_account_or_create(account) account = LedgerAccount( display_name=f"{user.uuid} Wallet", @@ -572,7 +578,7 @@ def setup_accounts( reference_uuid=user.uuid, currency="test", ) - lm.get_account_or_create(account=account) + ledger_manager.get_account_or_create(account=account) return _inner @@ -583,7 +589,7 @@ def session_with_tx_factory( session_manager: SessionManager, wall_manager: WallManager, utc_hour_ago: datetime, - thl_lm: ThlLedgerManager, + thl_ledger_manager: ThlLedgerManager, ) -> Callable[..., Session]: from generalresearch.models.thl.session import ( @@ -624,14 +630,16 @@ def session_with_tx_factory( status_code_1=status_code_1, ) - thl_lm.create_tx_task_complete( + thl_ledger_manager.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) + thl_ledger_manager.create_tx_bp_payment( + session=s, created=last_wall.finished, force=True + ) return s @@ -642,7 +650,7 @@ def session_with_tx_factory( def adj_to_fail_with_tx_factory( session_manager: SessionManager, wall_manager: WallManager, - thl_lm: ThlLedgerManager, + thl_ledger_manager: ThlLedgerManager, ) -> Callable[..., None]: from datetime import timedelta @@ -675,7 +683,7 @@ def adj_to_fail_with_tx_factory( adjusted_timestamp=created, ) - thl_lm.create_tx_task_adjustment( + thl_ledger_manager.create_tx_task_adjustment( wall=w1, user=session.user, created=created + timedelta(milliseconds=1), @@ -684,7 +692,7 @@ def adj_to_fail_with_tx_factory( session.wall_events = wall_manager.get_wall_events(session_id=session.id) session_manager.adjust_status(session=session) - thl_lm.create_tx_bp_adjustment( + thl_ledger_manager.create_tx_bp_adjustment( session=session, created=created + timedelta(milliseconds=2) ) @@ -695,7 +703,7 @@ def adj_to_fail_with_tx_factory( def adj_to_complete_with_tx_factory( session_manager: SessionManager, wall_manager: WallManager, - thl_lm: ThlLedgerManager, + thl_ledger_manager: ThlLedgerManager, ) -> Callable[..., None]: from datetime import timedelta @@ -714,7 +722,7 @@ def adj_to_complete_with_tx_factory( adjusted_timestamp=created, ) - thl_lm.create_tx_task_adjustment( + thl_ledger_manager.create_tx_task_adjustment( wall=w1, user=session.user, created=created + timedelta(milliseconds=1), @@ -723,7 +731,7 @@ def adj_to_complete_with_tx_factory( session.wall_events = wall_manager.get_wall_events(session_id=session.id) session_manager.adjust_status(session=session) - thl_lm.create_tx_bp_adjustment( + thl_ledger_manager.create_tx_bp_adjustment( session=session, created=created + timedelta(milliseconds=2) ) diff --git a/tests/incite/collections/test_df_collection_base.py b/tests/incite/collections/test_df_collection_base.py index e20b44b..6d715fa 100644 --- a/tests/incite/collections/test_df_collection_base.py +++ b/tests/incite/collections/test_df_collection_base.py @@ -5,7 +5,7 @@ import pandas as pd import pytest from pandera.pandas import DataFrameSchema -from generalresearch.incite.collections import ( +from generalresearch.incite.collections.base import ( DFCollection, DFCollectionType, ) @@ -53,7 +53,7 @@ class TestDFCollectionBaseProperties: data_type=df_coll_type, start=datetime(year=1800, month=1, day=1, tzinfo=UTC), finished=datetime(year=1900, month=1, day=1, tzinfo=UTC), - offset="100d", + offset="100D", archive_path=mnt_filepath.archive_path(enum_type=df_coll_type), ) @@ -67,7 +67,7 @@ class TestDFCollectionBaseProperties: data_type=df_coll_type, start=datetime(year=1800, month=1, day=1, tzinfo=UTC), finished=datetime(year=1900, month=1, day=1, tzinfo=UTC), - offset="100d", + offset="100D", archive_path=mnt_filepath.archive_path(enum_type=df_coll_type), ) diff --git a/tests/incite/collections/test_df_collection_item_base.py b/tests/incite/collections/test_df_collection_item_base.py index fd70bf0..83d4973 100644 --- a/tests/incite/collections/test_df_collection_item_base.py +++ b/tests/incite/collections/test_df_collection_item_base.py @@ -25,7 +25,7 @@ class TestDFCollectionItemBase: def test_init(self, mnt_filepath: GRLDatasets, df_coll_type: DFCollectionType): collection = DFCollection( data_type=df_coll_type, - offset="100d", + offset="100D", start=datetime(year=1800, month=1, day=1, tzinfo=UTC), finished=datetime(year=1900, month=1, day=1, tzinfo=UTC), archive_path=mnt_filepath.archive_path(enum_type=df_coll_type), @@ -53,7 +53,7 @@ class TestDFCollectionItemMethods: ): collection = DFCollection( data_type=df_coll_type, - offset="100d", + offset="100D", start=datetime(year=1800, month=1, day=1, tzinfo=UTC), finished=datetime(year=1900, month=1, day=1, tzinfo=UTC), archive_path=mnt_filepath.archive_path(enum_type=df_coll_type), @@ -70,7 +70,7 @@ class TestDFCollectionItemMethods: ): collection = DFCollection( data_type=df_coll_type, - offset="100d", + offset="100D", start=datetime(year=1800, month=1, day=1, tzinfo=UTC), finished=datetime(year=1900, month=1, day=1, tzinfo=UTC), archive_path=mnt_filepath.archive_path(enum_type=df_coll_type), diff --git a/tests/incite/test_interval_idx.py b/tests/incite/test_interval_idx.py index 03d29ea..04d0bb2 100644 --- a/tests/incite/test_interval_idx.py +++ b/tests/incite/test_interval_idx.py @@ -18,7 +18,7 @@ class TestIntervalIndex: # If the offset is longer than the end - start it will not # error. It will simply have 0 rows. iv_r: pd.IntervalIndex = pd.interval_range( - start=start, end=end, freq="30d", closed="left" + start=start, end=end, freq="30D", closed="left" ) assert isinstance(iv_r, pd.IntervalIndex) assert len(iv_r.to_list()) == 0 diff --git a/tests/managers/gr/test_business.py b/tests/managers/gr/test_business.py index 1a5d4fa..35c471e 100644 --- a/tests/managers/gr/test_business.py +++ b/tests/managers/gr/test_business.py @@ -7,8 +7,8 @@ from generalresearch.models.gr.business import ( Business, BusinessAddress, BusinessBankAccount, - TransferMethod, ) +from generalresearch.models.gr.definitions import TransferMethod if TYPE_CHECKING: from generalresearch.managers.gr.business import ( @@ -32,12 +32,12 @@ class TestBusinessBankAccountManager: def test_create( self, - business: Business, + gr_business: Business, business_bank_account_manager: BusinessBankAccountManager, ): instance = business_bank_account_manager.create( - business_id=business.id, + business_id=gr_business.id, uuid=uuid4().hex, transfer_method=TransferMethod.ACH, ) @@ -56,10 +56,12 @@ class TestBusinessBankAccountManager: class TestBusinessAddressManager: def test_create( - self, business: Business, business_address_manager: BusinessAddressManager + self, gr_business: Business, business_address_manager: BusinessAddressManager ): - res = business_address_manager.create(uuid=uuid4().hex, business_id=business.id) + res = business_address_manager.create( + uuid=uuid4().hex, business_id=gr_business.id + ) assert isinstance(res, BusinessAddress) assert isinstance(res.id, int) @@ -140,18 +142,20 @@ class TestBusinessManager: def test_get_uuids_by_user_id(self): pass - def test_get_by_uuid(self, business: Business, business_manager: BusinessManager): - instance = business_manager.get_by_uuid(business_uuid=business.uuid) + def test_get_by_uuid( + self, gr_business: Business, business_manager: BusinessManager + ): + instance = business_manager.get_by_uuid(business_uuid=gr_business.uuid) assert isinstance(instance, Business) - assert business.id == instance.id + assert gr_business.id == instance.id - def test_get_by_id(self, business: Business, business_manager: BusinessManager): - instance = business_manager.get_by_id(business_id=business.id) + def test_get_by_id(self, gr_business: Business, business_manager: BusinessManager): + instance = business_manager.get_by_id(business_id=gr_business.id) assert isinstance(instance, Business) - assert business.uuid == instance.uuid + assert gr_business.uuid == instance.uuid - def test_cache_key(self, business: Business): - assert "business:" in business.cache_key + def test_cache_key(self, gr_business: Business): + assert "business:" in gr_business.cache_key # def test_create_raise_on_duplicate(self): # b_uuid = uuid4().hex @@ -160,7 +164,7 @@ class TestBusinessManager: # business = BusinessManager.create( # uuid=b_uuid, # name=f"test-{b_uuid[:6]}") - # assert isinstance(business: Business, Business) + # assert isinstance(gr_business: Business, Business) # # # Try to make it again # with pytest.raises(expected_exception=psycopg.errors.UniqueViolation): diff --git a/tests/managers/thl/test_ledger/test_lm_accounts.py b/tests/managers/thl/test_ledger/test_lm_accounts.py index f5ed883..cdef99a 100644 --- a/tests/managers/thl/test_ledger/test_lm_accounts.py +++ b/tests/managers/thl/test_ledger/test_lm_accounts.py @@ -44,7 +44,7 @@ class TestLedgerAccountManagerNoResults: currency: LedgerCurrency, kind: str, acct_id: UUIDStr, - lm: LedgerManager, + ledger_manager: LedgerManager, ): """Try to query for accounts that we know don't exist and confirm that we either get the expected None result or it raises the correct @@ -54,40 +54,50 @@ class TestLedgerAccountManagerNoResults: # (1) .get_account is just a wrapper for .get_account_many_ but # call it either way - assert lm.get_account(qualified_name=qn, raise_on_error=False) is None + assert ( + ledger_manager.get_account(qualified_name=qn, raise_on_error=False) is None + ) with pytest.raises(expected_exception=LedgerAccountDoesntExistError): - lm.get_account(qualified_name=qn, raise_on_error=True) + ledger_manager.get_account(qualified_name=qn, raise_on_error=True) # (2) .get_account_if_exists is another wrapper - assert lm.get_account(qualified_name=qn, raise_on_error=False) is None + assert ( + ledger_manager.get_account(qualified_name=qn, raise_on_error=False) is None + ) def test_get_account_no_results_many( self, currency: LedgerCurrency, kind: str, acct_id: UUIDStr, - lm: LedgerManager, + ledger_manager: LedgerManager, ): qn = f"{currency}:{kind}:{acct_id}" # (1) .get_many_ - assert lm.get_account_many_(qualified_names=[qn], raise_on_error=False) == [] + assert ( + ledger_manager.get_account_many_(qualified_names=[qn], raise_on_error=False) + == [] + ) with pytest.raises(expected_exception=LedgerAccountDoesntExistError): - lm.get_account_many_(qualified_names=[qn], raise_on_error=True) + ledger_manager.get_account_many_(qualified_names=[qn], raise_on_error=True) # (2) .get_many - assert lm.get_account_many(qualified_names=[qn], raise_on_error=False) == [] + assert ( + ledger_manager.get_account_many(qualified_names=[qn], raise_on_error=False) + == [] + ) with pytest.raises(expected_exception=LedgerAccountDoesntExistError): - lm.get_account_many(qualified_names=[qn], raise_on_error=True) + ledger_manager.get_account_many(qualified_names=[qn], raise_on_error=True) # (3) .get_accounts(..) - assert lm.get_accounts_if_exists(qualified_names=[qn]) == [] + assert ledger_manager.get_accounts_if_exists(qualified_names=[qn]) == [] with pytest.raises(expected_exception=LedgerAccountDoesntExistError): - lm.get_accounts(qualified_names=[qn]) + ledger_manager.get_accounts(qualified_names=[qn]) @pytest.mark.parametrize( @@ -107,7 +117,7 @@ class TestLedgerAccountManagerCreate: currency: LedgerCurrency, account_type: AccountType, direction: Direction, - lm: LedgerManager, + ledger_manager: LedgerManager, ): """Confirm that the Permission values that are set on the Ledger Manger allow the Creation action to occur. @@ -124,11 +134,11 @@ class TestLedgerAccountManagerCreate: # (1) With no Permissions defined test_lm = LedgerManager( - pg_config=lm.pg_config, + pg_config=ledger_manager.pg_config, permissions=[], - redis_config=lm.redis_config, - cache_prefix=lm.cache_prefix, - testing=lm.testing, + redis_config=ledger_manager.redis_config, + cache_prefix=ledger_manager.cache_prefix, + testing=ledger_manager.testing, ) with pytest.raises(expected_exception=AssertionError) as excinfo: @@ -139,11 +149,11 @@ class TestLedgerAccountManagerCreate: # (2) With Permissions defined, but not CREATE test_lm = LedgerManager( - pg_config=lm.pg_config, + pg_config=ledger_manager.pg_config, permissions=[Permission.READ, Permission.UPDATE, Permission.DELETE], - redis_config=lm.redis_config, - cache_prefix=lm.cache_prefix, - testing=lm.testing, + redis_config=ledger_manager.redis_config, + cache_prefix=ledger_manager.cache_prefix, + testing=ledger_manager.testing, ) with pytest.raises(expected_exception=AssertionError) as excinfo: @@ -157,7 +167,7 @@ class TestLedgerAccountManagerCreate: currency: LedgerCurrency, account_type: AccountType, direction: Direction, - lm: LedgerManager, + ledger_manager: LedgerManager, ): """Confirm that the Permission values that are set on the Ledger Manger allow the Creation action to occur. @@ -174,11 +184,11 @@ class TestLedgerAccountManagerCreate: account_type=account_type, normal_balance=direction, ) - account = lm.create_account(account=acct_model) + account = ledger_manager.create_account(account=acct_model) assert isinstance(account, LedgerAccount) # Query for, and make sure the Account was saved in the DB - res = lm.get_account(qualified_name=qn, raise_on_error=True) + res = ledger_manager.get_account(qualified_name=qn, raise_on_error=True) assert res is not None assert account.uuid == res.uuid @@ -187,7 +197,7 @@ class TestLedgerAccountManagerCreate: currency: LedgerCurrency, account_type: AccountType, direction: Direction, - lm: LedgerManager, + ledger_manager: LedgerManager, ): """Confirm that the Permission values that are set on the Ledger Manger allow the Creation action to occur. @@ -204,27 +214,31 @@ class TestLedgerAccountManagerCreate: account_type=account_type, normal_balance=direction, ) - account = lm.get_account_or_create(account=acct_model) + account = ledger_manager.get_account_or_create(account=acct_model) assert isinstance(account, LedgerAccount) # Query for, and make sure the Account was saved in the DB - res = lm.get_account(qualified_name=qn, raise_on_error=True) + res = ledger_manager.get_account(qualified_name=qn, raise_on_error=True) assert res is not None assert account.uuid == res.uuid class TestLedgerAccountManagerGet: - def test_get(self, ledger_account: LedgerAccount, lm: LedgerManager): - res = lm.get_account(qualified_name=ledger_account.qualified_name) + def test_get(self, ledger_account: LedgerAccount, ledger_manager: LedgerManager): + res = ledger_manager.get_account(qualified_name=ledger_account.qualified_name) assert res is not None assert res.uuid == ledger_account.uuid - res = lm.get_account_many(qualified_names=[ledger_account.qualified_name]) + res = ledger_manager.get_account_many( + qualified_names=[ledger_account.qualified_name] + ) assert len(res) == 1 assert res[0].uuid == ledger_account.uuid - res = lm.get_accounts(qualified_names=[ledger_account.qualified_name]) + res = ledger_manager.get_accounts( + qualified_names=[ledger_account.qualified_name] + ) assert len(res) == 1 assert res[0].uuid == ledger_account.uuid @@ -237,15 +251,15 @@ class TestLedgerAccountManagerGet: ledger_account_credit: LedgerAccount, ledger_account_debit: LedgerAccount, ledger_tx: LedgerTransaction, - lm: LedgerManager, + ledger_manager: LedgerManager, ): - res = lm.get_account_balance(account=ledger_account) + res = ledger_manager.get_account_balance(account=ledger_account) assert res == 0 - res = lm.get_account_balance(account=ledger_account_credit) + res = ledger_manager.get_account_balance(account=ledger_account_credit) assert res == 100 - res = lm.get_account_balance(account=ledger_account_debit) + res = ledger_manager.get_account_balance(account=ledger_account_debit) assert res == 100 @pytest.mark.parametrize("n_times", range(5)) @@ -256,7 +270,7 @@ class TestLedgerAccountManagerGet: ledger_account_debit: LedgerAccount, ledger_tx: LedgerTransaction, n_times: PositiveInt, - lm: LedgerManager, + ledger_manager: LedgerManager, ): """Try searching for random metadata and confirm it's always 0 because Tx can be found. @@ -265,7 +279,7 @@ class TestLedgerAccountManagerGet: rand_value = uuid4().hex assert ( - lm.get_account_filtered_balance( + ledger_manager.get_account_filtered_balance( account=ledger_account, metadata_key=rand_key, metadata_value=rand_value ) == 0 @@ -275,7 +289,7 @@ class TestLedgerAccountManagerGet: # and that we can filter it back rand_amount = randint(10, 1_000) - lm.create_tx( + ledger_manager.create_tx( entries=[ LedgerEntry( direction=Direction.CREDIT, @@ -292,7 +306,7 @@ class TestLedgerAccountManagerGet: ) assert ( - lm.get_account_filtered_balance( + ledger_manager.get_account_filtered_balance( account=ledger_account_credit, metadata_key=rand_key, metadata_value=rand_value, @@ -301,7 +315,7 @@ class TestLedgerAccountManagerGet: ) assert ( - lm.get_account_filtered_balance( + ledger_manager.get_account_filtered_balance( account=ledger_account_debit, metadata_key=rand_key, metadata_value=rand_value, @@ -310,7 +324,7 @@ class TestLedgerAccountManagerGet: ) def test_get_balance_timerange_empty( - self, ledger_account: LedgerAccount, lm: LedgerManager + self, ledger_account: LedgerAccount, ledger_manager: LedgerManager ): - res = lm.get_account_balance_timerange(account=ledger_account) + res = ledger_manager.get_account_balance_timerange(account=ledger_account) assert res == 0 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 2e4ab5e..b0484ae 100644 --- a/tests/managers/thl/test_ledger/test_thl_lm_tx.py +++ b/tests/managers/thl/test_ledger/test_thl_lm_tx.py @@ -311,15 +311,14 @@ class TestThlLedgerTxManager: def test_create_tx_bp_payout_( self, product: Product, - thl_lm: ThlLedgerManager, - ledger_manager: LedgerManager, + thl_ledger_manager: ThlLedgerManager, currency: LedgerCurrency, ): rand_amount: USDCent = USDCent(randint(100, 1_000)) payoutevent_uuid = uuid4().hex # Create a BP Payout for a Product without any activity. - tx = thl_lm.create_tx_bp_payout_( + tx = thl_ledger_manager.create_tx_bp_payout_( product=product, amount=rand_amount, payoutevent_uuid=payoutevent_uuid, diff --git a/tests/managers/thl/test_payout.py b/tests/managers/thl/test_payout.py index 2494de8..ad101a4 100644 --- a/tests/managers/thl/test_payout.py +++ b/tests/managers/thl/test_payout.py @@ -86,7 +86,6 @@ class TestPayout: self, user: User, user_payout_event_manager: UserPayoutEventManager, - ledger_manager: LedgerManager, thl_ledger_manager: ThlLedgerManager, utc_now: datetime, ): @@ -128,11 +127,11 @@ class TestPayout: self, thl_web_rw: PostgresConfig, product: Product, - thl_lm: ThlLedgerManager, + thl_ledger_manager: ThlLedgerManager, brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager, utc_now: datetime, ) -> BrokerageProductPayoutEvent: - account = thl_lm.get_account_or_create_bp_wallet(product=product) + account = thl_ledger_manager.get_account_or_create_bp_wallet(product=product) bp_pe = BrokerageProductPayoutEvent( product_id=product.uuid, amount=USDCent(100), @@ -161,15 +160,14 @@ class TestPayout: self, product: Product, brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager, - thl_lm: ThlLedgerManager, - ledger_manager: LedgerManager, + thl_ledger_manager: ThlLedgerManager, utc_now: datetime, pending_bp_pe: BrokerageProductPayoutEvent, ): - thl_lm.get_account_or_create_bp_wallet(product=product) + thl_ledger_manager.get_account_or_create_bp_wallet(product=product) brokerage_product_payout_event_manager.create_tx_bp_payout_from_payout_event( - thl_ledger_manager=thl_lm, + thl_ledger_manager=thl_ledger_manager, bp_pe=pending_bp_pe, product=product, created=utc_now, @@ -177,7 +175,7 @@ class TestPayout: with pytest.raises(ValueError) as cm: brokerage_product_payout_event_manager.create_tx_bp_payout_from_payout_event( - thl_ledger_manager=thl_lm, + thl_ledger_manager=thl_ledger_manager, product=product, bp_pe=pending_bp_pe, created=utc_now, @@ -187,7 +185,6 @@ class TestPayout: def test_filter( self, thl_ledger_manager: ThlLedgerManager, - ledger_manager: LedgerManager, product: Product, user: User, user_payout_event_manager: UserPayoutEventManager, @@ -280,19 +277,18 @@ class TestBusinessPayoutEventManager: def test_base( self, - brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager, business_payout_event_manager: BusinessPayoutEventManager, delete_ledger_db: Callable[..., None], create_main_accounts: Callable[..., None], thl_ledger_manager: ThlLedgerManager, product_factory: Callable[..., Product], bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], - business: Business, + gr_business: Business, ): delete_ledger_db() create_main_accounts() - p1: Product = product_factory(business=business) + p1: Product = product_factory(business=gr_business) thl_ledger_manager.get_account_or_create_bp_wallet(product=p1) ach_id1 = uuid4().hex @@ -310,23 +306,25 @@ class TestBusinessPayoutEventManager: bp_payout_factory(product=p1, amount=USDCent(50), ext_ref_id=ach_id2) - business.prebuild_payouts( + gr_business.prebuild_payouts( bpem=business_payout_event_manager, ) - assert isinstance(business.payouts, list) - assert len(business.payouts) == 3 - assert business.payouts_total == sum([pe.amount for pe in business.payouts]) - assert business.payouts[0].created > business.payouts[1].created - assert len(business.payouts[0].bp_payouts) == 1 + assert isinstance(gr_business.payouts, list) + assert len(gr_business.payouts) == 3 + assert gr_business.payouts_total == sum( + [pe.amount for pe in gr_business.payouts] + ) + assert gr_business.payouts[0].created > gr_business.payouts[1].created + assert len(gr_business.payouts[0].bp_payouts) == 1 # Cannot pay out the same product twice in the same business payout # assert len(business.payouts[1].bp_payouts) == 2 - assert len(business.payouts[1].bp_payouts) == 1 + assert len(gr_business.payouts[1].bp_payouts) == 1 - assert business.payouts[0].ext_ref_id == ach_id2 - assert business.payouts[1].ext_ref_id == ach_id1 - assert business.payouts[2].ext_ref_id == "none" + assert gr_business.payouts[0].ext_ref_id == ach_id2 + assert gr_business.payouts[1].ext_ref_id == ach_id1 + assert gr_business.payouts[2].ext_ref_id == "none" def test_update_ext_reference_ids( self, @@ -345,13 +343,13 @@ class TestBusinessPayoutEventManager: mnt_filepath: GRLDatasets, product_manager: ProductManager, start: datetime, - business: Business, + gr_business: Business, ): delete_ledger_db() create_main_accounts() delete_df_collection(coll=ledger_collection) - p1: Product = product_factory(business=business) + p1: Product = product_factory(business=gr_business) u1: User = user_factory(product=p1) thl_ledger_manager.get_account_or_create_bp_wallet(product=p1) @@ -377,7 +375,7 @@ class TestBusinessPayoutEventManager: # We must build the balance to issue ACH/Wire ledger_collection.initial_load(client=None, sync=True) pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) - business.prebuild_balance( + gr_business.prebuild_balance( thl_pg_config=thl_web_rr, lm=thl_ledger_manager, ds=mnt_filepath, @@ -386,7 +384,7 @@ class TestBusinessPayoutEventManager: ) res = business_payout_event_manager.create_from_ach_or_wire( - business=business, + business=gr_business, amount=USDCent(100_01), pm=product_manager, thl_lm=thl_ledger_manager, @@ -558,7 +556,7 @@ class TestBusinessPayoutEventManager: create_main_accounts: Callable[..., None], delete_df_collection: Callable[..., None], ledger_collection: LedgerDFCollection, - business: Business, + gr_business: Business, user_factory: Callable[..., User], product_factory: Callable[..., Product], session_with_tx_factory: Callable[..., Session], @@ -581,7 +579,7 @@ class TestBusinessPayoutEventManager: create_main_accounts() delete_df_collection(coll=ledger_collection) - p1: Product = product_factory(business=business) + p1: Product = product_factory(business=gr_business) u1: User = user_factory(product=p1) thl_ledger_manager.get_account_or_create_bp_wallet(product=p1) @@ -603,7 +601,7 @@ class TestBusinessPayoutEventManager: ledger_collection.initial_load(client=None, sync=True) pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) - business.prebuild_balance( + gr_business.prebuild_balance( thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, @@ -613,7 +611,7 @@ class TestBusinessPayoutEventManager: with pytest.raises(expected_exception=AssertionError) as cm: business_payout_event_manager.create_from_ach_or_wire( - business=business, + business=gr_business, amount=USDCent(500), pm=product_manager, thl_lm=thl_ledger_manager, @@ -631,7 +629,7 @@ class TestBusinessPayoutEventManager: create_main_accounts: Callable[..., None], delete_df_collection: Callable[..., None], ledger_collection: LedgerDFCollection, - business: Business, + gr_business: Business, user_factory: Callable[..., User], product_factory: Callable[..., Product], session_with_tx_factory: Callable[..., None], @@ -648,9 +646,9 @@ class TestBusinessPayoutEventManager: create_main_accounts() delete_df_collection(coll=ledger_collection) - p1: Product = product_factory(business=business) - p2: Product = product_factory(business=business) - p3: Product = product_factory(business=business) + p1: Product = product_factory(business=gr_business) + p2: Product = product_factory(business=gr_business) + p3: Product = product_factory(business=gr_business) _: User = user_factory(product=p1) u2: User = user_factory(product=p2) u3: User = user_factory(product=p3) @@ -679,7 +677,7 @@ class TestBusinessPayoutEventManager: ledger_collection.initial_load(client=None, sync=True) pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) - business.prebuild_balance( + gr_business.prebuild_balance( thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, @@ -687,13 +685,13 @@ class TestBusinessPayoutEventManager: pop_ledger=pop_ledger_merge, ) - bb = business.balance + bb = gr_business.balance assert isinstance(bb, BusinessBalances) assert bb.payout == 475_00 # $500 * .95% = $475 assert bb.net == 475_00 bp1 = business_payout_event_manager.create_from_ach_or_wire( - business=business, + business=gr_business, amount=USDCent(100_00), pm=product_manager, thl_lm=thl_ledger_manager, @@ -705,7 +703,7 @@ class TestBusinessPayoutEventManager: assert len(bp1.bp_payouts) == 2 bp2 = business_payout_event_manager.create_from_ach_or_wire( - business=business, + business=gr_business, amount=USDCent(bb.available_balance), pm=product_manager, thl_lm=thl_ledger_manager, @@ -743,7 +741,7 @@ class TestBusinessPayoutEventManager: create_main_accounts: Callable[..., None], delete_df_collection: Callable[..., None], ledger_collection: LedgerDFCollection, - business: Business, + gr_business: Business, user_factory: Callable[..., User], product_factory: Callable[..., Product], session_with_tx_factory: Callable[..., None], @@ -768,9 +766,9 @@ class TestBusinessPayoutEventManager: create_main_accounts() delete_df_collection(coll=ledger_collection) - p1: Product = product_factory(business=business) - p2: Product = product_factory(business=business) - p3: Product = product_factory(business=business) + p1: Product = product_factory(business=gr_business) + p2: Product = product_factory(business=gr_business) + p3: Product = product_factory(business=gr_business) u1: User = user_factory(product=p1) u2: User = user_factory(product=p2) u3: User = user_factory(product=p3) @@ -813,10 +811,10 @@ class TestBusinessPayoutEventManager: started=start + timedelta(days=1, hours=3, minutes=1 + idx), ) - # Now that we paid out the business: Business, let's confirm the updated balances + # Now that we paid out the gr_business: Business, let's confirm the updated balances ledger_collection.initial_load(client=None, sync=True) pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) - business.prebuild_balance( + gr_business.prebuild_balance( thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, @@ -824,7 +822,7 @@ class TestBusinessPayoutEventManager: pop_ledger=pop_ledger_merge, ) - bb1 = business.balance + bb1 = gr_business.balance assert isinstance(bb1, BusinessBalances) pb1 = bb1.product_balances[0] pb2 = bb1.product_balances[1] @@ -848,18 +846,18 @@ class TestBusinessPayoutEventManager: assert pb2.recoup_usd_str == "$0.00" assert pb3.recoup_usd_str == "$0.00" - assert business.payouts is None - business.prebuild_payouts( + assert gr_business.payouts is None + gr_business.prebuild_payouts( thl_pg_config=thl_web_rr, thl_lm=thl_ledger_manager, bpem=business_payout_event_manager, ) - assert isinstance(business.payouts, list) - assert len(business.payouts) == 1 - assert business.payouts[0].ext_ref_id == ach_id1 + assert isinstance(gr_business.payouts, list) + assert len(gr_business.payouts) == 1 + assert gr_business.payouts[0].ext_ref_id == ach_id1 bp1 = business_payout_event_manager.create_from_ach_or_wire( - business=business, + business=gr_business, amount=USDCent(bb1.available_balance), pm=product_manager, thl_lm=thl_ledger_manager, @@ -937,7 +935,7 @@ class TestBusinessPayoutEventManager: create_main_accounts: Callable[..., None], delete_df_collection: Callable[..., None], ledger_collection: LedgerDFCollection, - business: Business, + gr_business: Business, user_factory: Callable[..., User], product_factory: Callable[..., Product], session_with_tx_factory: Callable[..., None], @@ -950,7 +948,7 @@ class TestBusinessPayoutEventManager: rm_pop_ledger_merge: Callable[..., None], ): """There are valid instances when we want issue a ACH or Wire to a - business: Business, but not for the full Available Balance amount in their + gr_business: Business, but not for the full Available Balance amount in their account. To test this, we'll create a Business with multiple Products, and @@ -965,9 +963,9 @@ class TestBusinessPayoutEventManager: create_main_accounts() delete_df_collection(coll=ledger_collection) - p1: Product = product_factory(business=business) - p2: Product = product_factory(business=business) - p3: Product = product_factory(business=business) + p1: Product = product_factory(business=gr_business) + p2: Product = product_factory(business=gr_business) + p3: Product = product_factory(business=gr_business) u1: User = user_factory(product=p1) u2: User = user_factory(product=p2) u3: User = user_factory(product=p3) @@ -988,20 +986,20 @@ class TestBusinessPayoutEventManager: # Now that we paid out the business: Business, let's confirm the updated balances ledger_collection.initial_load(client=None, sync=True) pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) - business.prebuild_balance( + gr_business.prebuild_balance( thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, pop_ledger=pop_ledger_merge, ) - business.prebuild_payouts( + gr_business.prebuild_payouts( bpem=business_payout_event_manager, ) # Confirm the initial amounts. - assert len(business.payouts) == 0 - bb1 = business.balance + assert len(gr_business.payouts) == 0 + bb1 = gr_business.balance assert isinstance(bb1, BusinessBalances) assert bb1.payout == 3 * 5 * 4750 @@ -1015,16 +1013,16 @@ class TestBusinessPayoutEventManager: assert bb1.product_balances[x].balance == 5 * 4750 assert bb1.product_balances[x].available_balance_usd_str == "$178.13" - assert business.payouts_total_str == "$0.00" - assert isinstance(business.balance, BusinessBalances) - assert business.balance.payment_usd_str == "$0.00" - assert business.balance.available_balance_usd_str == "$534.39" + assert gr_business.payouts_total_str == "$0.00" + assert isinstance(gr_business.balance, BusinessBalances) + assert gr_business.balance.payment_usd_str == "$0.00" + assert gr_business.balance.available_balance_usd_str == "$534.39" # This is the important part, even those the Business has $534.39 # available to it, we are only trying to issue out a $250.00 ACH or # Wire to the Business bp1 = business_payout_event_manager.create_from_ach_or_wire( - business=business, + business=gr_business, amount=USDCent(250_00), pm=product_manager, thl_lm=thl_ledger_manager, @@ -1033,7 +1031,7 @@ class TestBusinessPayoutEventManager: assert isinstance(bp1, BusinessPayoutEvent) assert len(bp1.bp_payouts) == 3 - # Now that we paid out the business: Business, let's confirm the updated + # Now that we paid out the gr_business: Business, let's confirm the updated # balances. Clear and rebuild the parquet files. rm_ledger_collection() rm_pop_ledger_merge() @@ -1043,25 +1041,23 @@ class TestBusinessPayoutEventManager: # Now rebuild and confirm the payouts, balance.payment, and the # balance.available_balance are reflective of having a $250 ACH/Wire # sent to the Business - business.prebuild_balance( + gr_business.prebuild_balance( thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, pop_ledger=pop_ledger_merge, ) - business.prebuild_payouts( - thl_pg_config=thl_web_rr, - thl_lm=thl_ledger_manager, + gr_business.prebuild_payouts( bpem=business_payout_event_manager, ) - assert isinstance(business.payouts, list) - assert len(business.payouts) == 1 - assert len(business.payouts[0].bp_payouts) == 3 - assert business.payouts_total_str == "$250.00" - assert isinstance(business.balance, BusinessBalances) - assert business.balance.payment_usd_str == "$250.00" - assert business.balance.available_balance_usd_str == "$346.88" + assert isinstance(gr_business.payouts, list) + assert len(gr_business.payouts) == 1 + assert len(gr_business.payouts[0].bp_payouts) == 3 + assert gr_business.payouts_total_str == "$250.00" + assert isinstance(gr_business.balance, BusinessBalances) + assert gr_business.balance.payment_usd_str == "$250.00" + assert gr_business.balance.available_balance_usd_str == "$346.88" def test_ach_tx_id_reference( self, @@ -1074,7 +1070,7 @@ class TestBusinessPayoutEventManager: create_main_accounts: Callable[..., None], delete_df_collection: Callable[..., None], ledger_collection: LedgerDFCollection, - business: Business, + gr_business: Business, user_factory: Callable[..., User], product_factory: Callable[..., Product], session_with_tx_factory: Callable[..., Session], @@ -1092,9 +1088,9 @@ class TestBusinessPayoutEventManager: create_main_accounts() delete_df_collection(coll=ledger_collection) - p1: Product = product_factory(business=business) - p2: Product = product_factory(business=business) - p3: Product = product_factory(business=business) + p1: Product = product_factory(business=gr_business) + p2: Product = product_factory(business=gr_business) + p3: Product = product_factory(business=gr_business) u1: User = user_factory(product=p1) u2: User = user_factory(product=p2) u3: User = user_factory(product=p3) @@ -1118,7 +1114,7 @@ class TestBusinessPayoutEventManager: rm_pop_ledger_merge() ledger_collection.initial_load(client=None, sync=True) pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) - business.prebuild_balance( + gr_business.prebuild_balance( thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, @@ -1127,7 +1123,7 @@ class TestBusinessPayoutEventManager: ) bp1 = business_payout_event_manager.create_from_ach_or_wire( - business=business, + business=gr_business, amount=USDCent(100_01), transaction_id=ach_id1, pm=product_manager, @@ -1139,7 +1135,7 @@ class TestBusinessPayoutEventManager: rm_pop_ledger_merge() ledger_collection.initial_load(client=None, sync=True) pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) - business.prebuild_balance( + gr_business.prebuild_balance( thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, @@ -1148,7 +1144,7 @@ class TestBusinessPayoutEventManager: ) bp2 = business_payout_event_manager.create_from_ach_or_wire( - business=business, + business=gr_business, amount=USDCent(100_02), transaction_id=ach_id2, pm=product_manager, @@ -1163,18 +1159,18 @@ class TestBusinessPayoutEventManager: rm_pop_ledger_merge() ledger_collection.initial_load(client=None, sync=True) pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) - business.prebuild_payouts( + gr_business.prebuild_payouts( thl_pg_config=thl_web_rr, thl_lm=thl_ledger_manager, bpem=business_payout_event_manager, ) - business.prebuild_balance( + gr_business.prebuild_balance( thl_pg_config=thl_web_rr, lm=ledger_manager, ds=mnt_filepath, client=client_no_amm, pop_ledger=pop_ledger_merge, ) - assert isinstance(business.payouts, list) - assert business.payouts[0].ext_ref_id == ach_id2 - assert business.payouts[1].ext_ref_id == ach_id1 + assert isinstance(gr_business.payouts, list) + assert gr_business.payouts[0].ext_ref_id == ach_id2 + assert gr_business.payouts[1].ext_ref_id == ach_id1 diff --git a/tests/managers/thl/test_session_manager.py b/tests/managers/thl/test_session_manager.py index 30fd9ec..67a802e 100644 --- a/tests/managers/thl/test_session_manager.py +++ b/tests/managers/thl/test_session_manager.py @@ -137,19 +137,19 @@ class TestSessionManagerFilter: def test_business( self, product_factory: Callable[..., Product], - business: Business, + gr_business: Business, user_factory: Callable[..., User], session_manager: SessionManager, utc_hour_ago: datetime, thl_web_rr: PostgresConfig, ): - p1 = product_factory(business=business) + p1 = product_factory(business=gr_business) for _ in range(5): u = user_factory(product=p1) session_manager.create(started=utc_hour_ago, user=u, uuid_id=uuid4().hex) - business.prefetch_products(thl_pg_config=thl_web_rr) - assert len(business.product_uuids) == 1 - res = session_manager.filter(product_uuids=business.product_uuids) + gr_business.prefetch_products(thl_pg_config=thl_web_rr) + assert len(gr_business.product_uuids) == 1 + res = session_manager.filter(product_uuids=gr_business.product_uuids) assert len(res) == 5 diff --git a/tests/models/gr/test_authentication.py b/tests/models/gr/test_authentication.py index ac1298f..059a0a4 100644 --- a/tests/models/gr/test_authentication.py +++ b/tests/models/gr/test_authentication.py @@ -116,7 +116,7 @@ class TestGRUser: class TestGRUserMethods: - def test_cache_key(self, gr_user: GRUser, gr_redis: RedisConfig): + def test_cache_key(self, gr_user: GRUser): assert isinstance(gr_user.cache_key, str) assert ":" in gr_user.cache_key assert str(gr_user.id) in gr_user.cache_key @@ -124,13 +124,12 @@ class TestGRUserMethods: def test_to_redis( self, gr_user: GRUser, - gr_redis: Redis, team: Team, - business: Business, + gr_business: Business, product_factory: Callable[..., Product], membership_factory: Callable[..., Membership], ): - product_factory(team=team, business=business) + product_factory(team=team, business=gr_business) membership_factory(team=team, gr_user=gr_user) res = gr_user.to_redis() @@ -144,31 +143,30 @@ class TestGRUserMethods: def test_set_cache( self, gr_user: GRUser, - gr_user_token: GRToken, - gr_redis: Redis, gr_db: PostgresConfig, thl_web_rr: PostgresConfig, gr_redis_config: RedisConfig, ): - assert gr_redis.get(name=gr_user.cache_key) is None - assert gr_redis.get(name=f"{gr_user.cache_key}:team_uuids") is None - assert gr_redis.get(name=f"{gr_user.cache_key}:business_uuids") is None - assert gr_redis.get(name=f"{gr_user.cache_key}:product_uuids") is None + + client = gr_redis_config.create_redis_client() + + assert client.get(name=gr_user.cache_key) is None + assert client.get(name=f"{gr_user.cache_key}:team_uuids") is None + assert client.get(name=f"{gr_user.cache_key}:business_uuids") is None + assert client.get(name=f"{gr_user.cache_key}:product_uuids") is None gr_user.set_cache( pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config ) - assert gr_redis.get(name=gr_user.cache_key) is not None - assert gr_redis.get(name=f"{gr_user.cache_key}:team_uuids") is not None - assert gr_redis.get(name=f"{gr_user.cache_key}:business_uuids") is not None - assert gr_redis.get(name=f"{gr_user.cache_key}:product_uuids") is not None + assert client.get(name=gr_user.cache_key) is not None + assert client.get(name=f"{gr_user.cache_key}:team_uuids") is not None + assert client.get(name=f"{gr_user.cache_key}:business_uuids") is not None + assert client.get(name=f"{gr_user.cache_key}:product_uuids") is not None def test_set_cache_gr_user( self, gr_user: GRUser, - gr_user_token: GRToken, - gr_redis: RedisConfig, gr_redis_config: RedisConfig, gr_db: PostgresConfig, thl_web_rr: PostgresConfig, @@ -179,6 +177,8 @@ class TestGRUserMethods: ): from generalresearch.models.gr.authentication import GRUser + client = gr_redis_config.create_redis_client() + p1 = product_factory(team=team) membership_factory(team=team, gr_user=gr_user) @@ -186,7 +186,7 @@ class TestGRUserMethods: pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config ) - res: str = gr_redis.get(name=gr_user.cache_key) + res: str = client.get(name=gr_user.cache_key) gru2 = GRUser.from_redis(res) assert gr_user.model_dump_json( @@ -203,9 +203,6 @@ class TestGRUserMethods: def test_set_cache_team_uuids( self, gr_user: GRUser, - membership: Membership, - gr_user_token: GRToken, - gr_redis: Redis, gr_db: PostgresConfig, thl_web_rr: PostgresConfig, product_factory: Callable[..., Product], @@ -213,11 +210,12 @@ class TestGRUserMethods: gr_redis_config: RedisConfig, ): product_factory(team=team) + client = gr_redis_config.create_redis_client() gr_user.set_cache( pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config ) - res = json.loads(gr_redis.get(name=f"{gr_user.cache_key}:team_uuids")) + res = json.loads(client.get(name=f"{gr_user.cache_key}:team_uuids")) assert len(res) == 1 assert gr_user.team_uuids == res @@ -225,29 +223,27 @@ class TestGRUserMethods: def test_set_cache_business_uuids( self, gr_user: GRUser, - gr_redis: Redis, gr_db: PostgresConfig, thl_web_rr: PostgresConfig, product_factory: Callable[..., Product], - business: Business, + gr_business: Business, team: Team, gr_redis_config: RedisConfig, ): - product_factory(team=team, business=business) + product_factory(team=team, business=gr_business) gr_user.set_cache( pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config ) - res = json.loads(gr_redis.get(name=f"{gr_user.cache_key}:business_uuids")) + + client = gr_redis_config.create_redis_client() + res = json.loads(client.get(name=f"{gr_user.cache_key}:business_uuids")) assert len(res) == 1 assert gr_user.business_uuids == res def test_set_cache_product_uuids( self, gr_user: GRUser, - membership: Membership, - gr_user_token: GRToken, - gr_redis: Redis, gr_db: PostgresConfig, thl_web_rr: PostgresConfig, product_factory: Callable[..., Product], @@ -259,7 +255,8 @@ class TestGRUserMethods: gr_user.set_cache( pg_config=gr_db, thl_web_rr=thl_web_rr, redis_config=gr_redis_config ) - res = json.loads(gr_redis.get(name=f"{gr_user.cache_key}:product_uuids")) + client = gr_redis_config.create_redis_client() + res = json.loads(client.get(name=f"{gr_user.cache_key}:product_uuids")) assert len(res) == 1 assert gr_user.product_uuids == res diff --git a/tests/models/gr/test_business.py b/tests/models/gr/test_business.py index 2c12da1..90e69db 100644 --- a/tests/models/gr/test_business.py +++ b/tests/models/gr/test_business.py @@ -64,10 +64,8 @@ class TestBusinessBankAccount: gr_business: Business, business_bank_account_manager: BusinessBankAccountManager, ): - from generalresearch.models.gr.business import ( - BusinessBankAccount, - TransferMethod, - ) + from generalresearch.models.gr.business import BusinessBankAccount + from generalresearch.models.gr.definitions import TransferMethod instance = business_bank_account_manager.create( business_id=gr_business.id, @@ -115,7 +113,7 @@ class TestBusiness: @pytest.fixture def offset(self) -> str: - return "30d" + return "30D" @pytest.fixture def duration(self) -> timedelta | None: @@ -222,46 +220,46 @@ class TestBusiness: def test_teams( self, - business: Business, + gr_business: Business, team: Team, team_manager: TeamManager, gr_db: PostgresConfig, ): - assert business.teams is None + assert gr_business.teams is None - business.prefetch_teams(pg_config=gr_db) - assert isinstance(business.teams, list) - assert len(business.teams) == 0 + gr_business.prefetch_teams(pg_config=gr_db) + assert isinstance(gr_business.teams, list) + assert len(gr_business.teams) == 0 - team_manager.add_business(team=team, business=business) - assert len(business.teams) == 0 - business.prefetch_teams(pg_config=gr_db) - assert len(business.teams) == 1 + team_manager.add_business(team=team, business=gr_business) + assert len(gr_business.teams) == 0 + gr_business.prefetch_teams(pg_config=gr_db) + assert len(gr_business.teams) == 1 def test_products( self, - business: Business, + gr_business: Business, product_factory: Callable[..., Product], product_manager: ProductManager, ): - p1 = product_factory(business=business) - assert business.products is None + p1 = product_factory(business=gr_business) + assert gr_business.products is None - business.prefetch_products(product_manager=product_manager) - assert isinstance(business.products, list) - assert len(business.products) == 1 - assert isinstance(business.products[0], Product) + gr_business.prefetch_products(product_manager=product_manager) + assert isinstance(gr_business.products, list) + assert len(gr_business.products) == 1 + assert isinstance(gr_business.products[0], Product) - assert business.products[0].uuid == p1.uuid + assert gr_business.products[0].uuid == p1.uuid # Add two more, but list is still one until we prefetch - product_factory(business=business) - product_factory(business=business) - assert len(business.products) == 1 + product_factory(business=gr_business) + product_factory(business=gr_business) + assert len(gr_business.products) == 1 - business.prefetch_products(product_manager=product_manager) - assert len(business.products) == 3 + gr_business.prefetch_products(product_manager=product_manager) + assert len(gr_business.products) == 3 def test_bank_accounts( self, @@ -306,7 +304,6 @@ class TestBusiness: self, gr_business: Business, product_factory: Callable[..., Product], - thl_web_rr: PostgresConfig, thl_ledger_manager: ThlLedgerManager, business_payout_event_manager: BusinessPayoutEventManager, ): @@ -322,8 +319,6 @@ class TestBusiness: thl_ledger_manager.get_account_or_create_bp_wallet(product=p) gr_business.prebuild_payouts( - thl_pg_config=thl_web_rr, - thl_lm=thl_ledger_manager, bpem=business_payout_event_manager, ) assert isinstance(gr_business.payouts, list) @@ -335,7 +330,6 @@ class TestBusiness: product_factory: Callable[..., Product], bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], thl_ledger_manager: ThlLedgerManager, - thl_web_rr: PostgresConfig, business_payout_event_manager: BusinessPayoutEventManager, create_main_accounts: Callable[..., None], ): @@ -351,8 +345,6 @@ class TestBusiness: ) gr_business.prebuild_payouts( - thl_pg_config=thl_web_rr, - thl_lm=thl_ledger_manager, bpem=business_payout_event_manager, ) assert len(gr_business.payouts) == 1 @@ -478,7 +470,7 @@ class TestBusinessBalance: @pytest.fixture def offset(self) -> str: - return "30d" + return "30D" @pytest.fixture def duration(self) -> timedelta | None: @@ -1190,15 +1182,14 @@ class TestBusinessMethods: ) -> timedelta | None: return None - def test_cache_key(self, business: Business): - assert isinstance(business.cache_key, str) - assert ":" in business.cache_key - assert str(business.uuid) in business.cache_key + def test_cache_key(self, gr_business: Business): + assert isinstance(gr_business.cache_key, str) + assert ":" in gr_business.cache_key + assert str(gr_business.uuid) in gr_business.cache_key def test_set_cache( self, gr_business: Business, - gr_redis: RedisConfig, gr_db: PostgresConfig, thl_web_rr: PostgresConfig, client_no_amm: DaskClient, @@ -1218,7 +1209,8 @@ class TestBusinessMethods: gr_redis_config: RedisConfig, mnt_gr_api_dir: Path, ): - assert gr_redis.get(name=gr_business.cache_key) is None + client = gr_redis_config.create_redis_client() + assert client.get(name=gr_business.cache_key) is None p1 = product_factory(team=team, business=gr_business) u1 = user_factory(product=p1) @@ -1244,7 +1236,7 @@ class TestBusinessMethods: mnt_gr_api=mnt_gr_api_dir, ) - assert gr_redis.hgetall(name=gr_business.cache_key) is not None + assert client.hgetall(name=gr_business.cache_key) is not None from generalresearch.models.gr.business import Business # We're going to pull only a specific year, but make sure that @@ -1367,7 +1359,7 @@ class TestBusinessMethods: session_factory: Callable[..., Session], product_factory: Callable[..., Product], delete_df_collection: Callable[..., None], - business: Business, + gr_business: Business, mnt_filepath: GRLDatasets, mnt_gr_api_dir: Path, ): @@ -1375,8 +1367,8 @@ class TestBusinessMethods: delete_df_collection(coll=wall_collection) delete_df_collection(coll=session_collection) - p1 = product_factory(business=business) - p2 = product_factory(business=business) + p1 = product_factory(business=gr_business) + p2 = product_factory(business=gr_business) for p in [p1, p2]: u = user_factory(product=p) @@ -1397,7 +1389,7 @@ class TestBusinessMethods: pg_config=thl_web_rr, ) - business.prebuild_enriched_session_parquet( + gr_business.prebuild_enriched_session_parquet( thl_pg_config=thl_web_rr, ds=mnt_filepath, client=client_no_amm, @@ -1407,7 +1399,9 @@ class TestBusinessMethods: # Now try to read from path df = pd.read_parquet( - os.path.join(mnt_gr_api_dir, "pop_session", f"{business.file_key}.parquet") + os.path.join( + mnt_gr_api_dir, "pop_session", f"{gr_business.file_key}.parquet" + ) ) assert isinstance(df, pd.DataFrame) diff --git a/tests/models/gr/test_team.py b/tests/models/gr/test_team.py index c1ae6d6..aa2de45 100644 --- a/tests/models/gr/test_team.py +++ b/tests/models/gr/test_team.py @@ -152,7 +152,6 @@ class TestTeamMethods: def test_set_cache( self, team: Team, - gr_redis: RedisConfig, gr_db: PostgresConfig, thl_web_rr: PostgresConfig, gr_redis_config: RedisConfig, @@ -162,7 +161,8 @@ class TestTeamMethods: enriched_wall_merge: EnrichedWallMerge, enriched_session_merge: EnrichedSessionMerge, ): - assert gr_redis.get(name=team.cache_key) is None + client = gr_redis_config.create_redis_client() + assert client.get(name=team.cache_key) is None team.set_cache( pg_config=gr_db, @@ -175,7 +175,7 @@ class TestTeamMethods: enriched_session=enriched_session_merge, ) - assert gr_redis.hgetall(name=team.cache_key) is not None + assert client.hgetall(name=team.cache_key) is not None def test_set_cache_team( self, diff --git a/tests/models/test_finance.py b/tests/models/test_finance.py index eabc877..c579d78 100644 --- a/tests/models/test_finance.py +++ b/tests/models/test_finance.py @@ -760,7 +760,7 @@ class TestPOPFinancialData: duration: timedelta, create_main_accounts: Callable[..., None], session_with_tx_factory: Callable[..., Session], - thl_lm: ThlLedgerManager, + thl_ledger_manager: ThlLedgerManager, delete_df_collection: Callable[..., None], delete_ledger_db: Callable[..., None], ): @@ -798,8 +798,10 @@ class TestPOPFinancialData: last_item_finish = item_finishes[0] accounts = [] - for _ in users: - account = thl_lm.get_account_or_create_bp_wallet(product=u.product) + for _u in users: + account = thl_ledger_manager.get_account_or_create_bp_wallet( + product=_u.product + ) accounts.append(account) account_ids = [a.uuid for a in accounts] @@ -856,7 +858,7 @@ class TestBusinessBalanceData: user_factory: Callable[..., User], product: Product, create_main_accounts: Callable[..., None], - thl_lm: ThlLedgerManager, + thl_ledger_manager: ThlLedgerManager, thl_web_rr: PostgresConfig, delete_df_collection: Callable[..., None], delete_ledger_db: Callable[..., None], @@ -886,7 +888,9 @@ class TestBusinessBalanceData: pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) # assert pop_ledger_merge.progress.has_archive.eq(True).all() - account: LedgerAccount = thl_lm.get_account_or_create_bp_wallet(product=product) + account: LedgerAccount = thl_ledger_manager.get_account_or_create_bp_wallet( + product=product + ) ddf = pop_ledger_merge.ddf( force_rr_latest=False, diff --git a/tests/models/thl/test_payout.py b/tests/models/thl/test_payout.py index 927687e..cc00f33 100644 --- a/tests/models/thl/test_payout.py +++ b/tests/models/thl/test_payout.py @@ -10,8 +10,8 @@ from generalresearch.models.gr import Team from generalresearch.models.gr.business import ( Business, BusinessAddress, - BusinessType, ) +from generalresearch.models.gr.definitions import BusinessType from generalresearch.models.thl.payout import ( BrokerageProductPayoutEvent, BusinessPayoutEvent, diff --git a/tests/models/thl/test_product.py b/tests/models/thl/test_product.py index cc0fa8e..a1b3688 100644 --- a/tests/models/thl/test_product.py +++ b/tests/models/thl/test_product.py @@ -603,7 +603,7 @@ class TestProductFinancials: @pytest.fixture def offset(self) -> str: - return "30d" + return "30D" @pytest.fixture def duration(self) -> timedelta | None: @@ -611,12 +611,12 @@ class TestProductFinancials: def test_balance( self, - business: Business, + gr_business: Business, product_factory: Callable[..., Product], user_factory: Callable[..., User], mnt_filepath: GRLDatasets, bp_payout_factory: Callable[..., BrokerageProductPayoutEvent], - thl_lm: ThlLedgerManager, + thl_ledger_manager: ThlLedgerManager, start: datetime, brokerage_product_payout_event_manager: BrokerageProductPayoutEventManager, session_with_tx_factory: Callable[..., Session], @@ -633,33 +633,54 @@ class TestProductFinancials: from generalresearch.currency import USDCent - p1: Product = product_factory(business=business) + p1: Product = product_factory(business=gr_business) u1: User = user_factory(product=p1) - bp_wallet = thl_lm.get_account_or_create_bp_wallet(product=p1) - thl_lm.get_account_or_create_user_wallet(user=u1) + bp_wallet = thl_ledger_manager.get_account_or_create_bp_wallet(product=p1) + thl_ledger_manager.get_account_or_create_user_wallet(user=u1) brokerage_product_payout_event_manager.set_account_lookup_table(thl_lm=thl_lm) - assert len(thl_lm.get_tx_filtered_by_account(account_uuid=bp_wallet.uuid)) == 0 + assert ( + len( + thl_ledger_manager.get_tx_filtered_by_account( + account_uuid=bp_wallet.uuid + ) + ) + == 0 + ) session_with_tx_factory( user=u1, wall_req_cpi=Decimal(".50"), started=start + timedelta(days=1), ) - assert thl_lm.get_account_balance(account=bp_wallet) == 48 - assert len(thl_lm.get_tx_filtered_by_account(account_uuid=bp_wallet.uuid)) == 1 + assert thl_ledger_manager.get_account_balance(account=bp_wallet) == 48 + assert ( + len( + thl_ledger_manager.get_tx_filtered_by_account( + account_uuid=bp_wallet.uuid + ) + ) + == 1 + ) session_with_tx_factory( user=u1, wall_req_cpi=Decimal("1.00"), started=start + timedelta(days=2), ) - assert thl_lm.get_account_balance(account=bp_wallet) == 143 - assert len(thl_lm.get_tx_filtered_by_account(account_uuid=bp_wallet.uuid)) == 2 + assert thl_ledger_manager.get_account_balance(account=bp_wallet) == 143 + assert ( + len( + thl_ledger_manager.get_tx_filtered_by_account( + account_uuid=bp_wallet.uuid + ) + ) + == 2 + ) with pytest.raises(expected_exception=AssertionError) as cm: p1.prebuild_balance( - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, ds=mnt_filepath, client=client_no_amm, ) @@ -669,7 +690,7 @@ class TestProductFinancials: pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) p1.prebuild_balance( - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, ds=mnt_filepath, client=client_no_amm, ) @@ -683,7 +704,7 @@ class TestProductFinancials: assert p1.balance.available_balance == 108 p1.prebuild_payouts( - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, bp_pem=brokerage_product_payout_event_manager, ) assert p1.payouts is not None @@ -700,7 +721,14 @@ class TestProductFinancials: skip_wallet_balance_check=True, skip_one_per_day_check=True, ) - assert len(thl_lm.get_tx_filtered_by_account(account_uuid=bp_wallet.uuid)) == 3 + assert ( + len( + thl_ledger_manager.get_tx_filtered_by_account( + account_uuid=bp_wallet.uuid + ) + ) + == 3 + ) # RM the entire directories shutil.rmtree(ledger_collection.archive_path) @@ -712,7 +740,7 @@ class TestProductFinancials: pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) p1.prebuild_balance( - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, ds=mnt_filepath, client=client_no_amm, ) @@ -726,7 +754,7 @@ class TestProductFinancials: assert p1.balance.available_balance == 70 p1.prebuild_payouts( - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, bp_pem=brokerage_product_payout_event_manager, ) assert p1.payouts is not None @@ -743,7 +771,14 @@ class TestProductFinancials: skip_wallet_balance_check=True, skip_one_per_day_check=True, ) - assert len(thl_lm.get_tx_filtered_by_account(account_uuid=bp_wallet.uuid)) == 4 + assert ( + len( + thl_ledger_manager.get_tx_filtered_by_account( + account_uuid=bp_wallet.uuid + ) + ) + == 4 + ) # RM the entire directories shutil.rmtree(ledger_collection.archive_path) @@ -755,7 +790,7 @@ class TestProductFinancials: pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) p1.prebuild_balance( - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, ds=mnt_filepath, client=client_no_amm, ) @@ -769,7 +804,7 @@ class TestProductFinancials: assert p1.balance.available_balance == 66 p1.prebuild_payouts( - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, bp_pem=brokerage_product_payout_event_manager, ) assert p1.payouts is not None @@ -786,7 +821,7 @@ class TestProductBalance: @pytest.fixture def offset(self) -> str: - return "30d" + return "30D" @pytest.fixture def duration(self) -> timedelta | None: @@ -796,7 +831,7 @@ class TestProductBalance: self, product: Product, mnt_filepath: GRLDatasets, - thl_lm: ThlLedgerManager, + thl_ledger_manager: ThlLedgerManager, client_no_amm: DaskClient, delete_ledger_db: Callable[..., None], create_main_accounts: Callable[..., None], @@ -826,7 +861,7 @@ class TestProductBalance: pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) # 2. Payout and build Parquets 2nd time - payout_event_manager.set_account_lookup_table(thl_lm=thl_lm) + payout_event_manager.set_account_lookup_table(thl_lm=thl_ledger_manager) bp_payout_factory( product=product, amount=USDCent(71), @@ -840,7 +875,7 @@ class TestProductBalance: with pytest.raises(expected_exception=AssertionError) as cm: product.prebuild_balance( - thl_lm=thl_lm, ds=mnt_filepath, client=client_no_amm + thl_lm=thl_ledger_manager, ds=mnt_filepath, client=client_no_amm ) assert "Sql and Parquet Balance inconsistent" in str(cm) @@ -848,7 +883,7 @@ class TestProductBalance: self, product: Product, mnt_filepath: GRLDatasets, - thl_lm: ThlLedgerManager, + thl_ledger_manager: ThlLedgerManager, client_no_amm: DaskClient, delete_ledger_db: Callable[..., None], create_main_accounts: Callable[..., None], @@ -885,7 +920,7 @@ class TestProductBalance: # 2. Payout and build Parquets 2nd time but this payout is "now" # so it hasn't already been archived - payout_event_manager.set_account_lookup_table(thl_lm=thl_lm) + payout_event_manager.set_account_lookup_table(thl_lm=thl_ledger_manager) bp_payout_factory( product=product, amount=USDCent(71), @@ -898,7 +933,9 @@ class TestProductBalance: pop_ledger_merge.build(client=client_no_amm, ledger_coll=ledger_collection) # We just want to call this to confirm it doesn't raise. - product.prebuild_balance(thl_lm=thl_lm, ds=mnt_filepath, client=client_no_amm) + product.prebuild_balance( + thl_lm=thl_ledger_manager, ds=mnt_filepath, client=client_no_amm + ) class TestProductPOPFinancial: @@ -909,7 +946,7 @@ class TestProductPOPFinancial: @pytest.fixture def offset(self) -> str: - return "30d" + return "30D" @pytest.fixture def duration(self) -> timedelta | None: @@ -919,7 +956,7 @@ class TestProductPOPFinancial: self, product: Product, mnt_filepath: GRLDatasets, - thl_lm: ThlLedgerManager, + thl_ledger_manager: ThlLedgerManager, client_no_amm: DaskClient, delete_ledger_db: Callable[..., None], create_main_accounts: Callable[..., None], @@ -955,7 +992,7 @@ class TestProductPOPFinancial: # --- test --- assert product.pop_financial is None product.prebuild_pop_financial( - thl_lm=thl_lm, + thl_lm=thl_ledger_manager, ds=mnt_filepath, client=client_no_amm, pop_ledger=pop_ledger_merge, @@ -982,7 +1019,7 @@ class TestProductCache: @pytest.fixture def offset(self) -> str: - return "30d" + return "30D" @pytest.fixture def duration(self) -> timedelta | None: -- cgit v1.2.3