aboutsummaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorMax Nanis2026-08-27 17:46:11 -0700
committerMax Nanis2026-08-27 17:46:11 -0700
commitaeeb7fef2594ccd34fbe96a77f6c5b392299fed7 (patch)
treede0896615df0c9747208fc26eb1ddde60f248286
parentfdc170938ac4ac8fa5d9d4df1936e6dc0777f291 (diff)
downloadgeneralresearch-aeeb7fef2594ccd34fbe96a77f6c5b392299fed7.tar.gz
generalresearch-aeeb7fef2594ccd34fbe96a77f6c5b392299fed7.zip
Ruff afternoon
-rw-r--r--generalresearch/managers/thl/product.py1
-rw-r--r--generalresearch/models/network/nmap/parser.py8
-rw-r--r--generalresearch/models/precision/question.py4
-rw-r--r--generalresearch/models/precision/survey.py17
-rw-r--r--generalresearch/models/prodege/survey.py3
-rw-r--r--generalresearch/models/spectrum/question.py9
-rw-r--r--generalresearch/models/spectrum/survey.py13
-rw-r--r--generalresearch/models/spectrum/task_collection.py2
-rw-r--r--generalresearch/models/thl/contest/contest.py10
-rw-r--r--generalresearch/models/thl/contest/leaderboard.py13
-rw-r--r--generalresearch/models/thl/contest/milestone.py20
-rw-r--r--generalresearch/models/thl/contest/raffle.py46
-rw-r--r--generalresearch/models/thl/finance.py12
-rw-r--r--generalresearch/models/thl/product.py36
-rw-r--r--generalresearch/models/thl/profiling/other_option.py4
-rw-r--r--generalresearch/models/thl/profiling/upk_question.py13
-rw-r--r--generalresearch/models/thl/profiling/user_question_answer.py31
-rw-r--r--pyproject.toml5
-rw-r--r--tests/incite/collections/test_df_collection_item_thl_web.py273
-rw-r--r--tests/incite/mergers/foundations/test_enriched_session.py52
-rw-r--r--tests/incite/mergers/foundations/test_enriched_task_adjust.py38
-rw-r--r--tests/incite/mergers/foundations/test_enriched_wall.py73
-rw-r--r--tests/incite/mergers/test_merge_collection.py53
-rw-r--r--tests/incite/mergers/test_merge_collection_item.py25
-rw-r--r--tests/incite/mergers/test_pop_ledger.py109
-rw-r--r--tests/incite/mergers/test_ym_survey_merge.py55
-rw-r--r--tests/incite/schemas/test_admin_responses.py36
-rw-r--r--tests/incite/schemas/test_thl_web.py8
-rw-r--r--tests/incite/test_collection_base.py62
-rw-r--r--tests/incite/test_collection_base_item.py74
-rw-r--r--tests/incite/test_interval_idx.py6
-rw-r--r--tests/managers/gr/test_business.py60
-rw-r--r--tests/managers/gr/test_team.py58
-rw-r--r--tests/managers/network/test_label.py4
-rw-r--r--tests/managers/thl/test_contest/test_leaderboard.py5
-rw-r--r--tests/managers/thl/test_contest/test_milestone.py5
-rw-r--r--tests/managers/thl/test_contest/test_raffle.py12
-rw-r--r--tests/managers/thl/test_ledger/test_lm_accounts.py2
-rw-r--r--tests/managers/thl/test_ledger/test_lm_tx.py4
-rw-r--r--tests/managers/thl/test_ledger/test_thl_lm_accounts.py11
-rw-r--r--tests/managers/thl/test_ledger/test_thl_lm_bp_payout.py54
-rw-r--r--tests/managers/thl/test_ledger/test_thl_lm_tx.py26
-rw-r--r--tests/managers/thl/test_payout.py2
-rw-r--r--tests/managers/thl/test_survey.py6
-rw-r--r--tests/managers/thl/test_user_manager/test_base.py2
-rw-r--r--tests/models/network/test_nmap.py2
-rw-r--r--tests/models/spectrum/test_survey.py44
-rw-r--r--tests/models/thl/test_product.py4
48 files changed, 815 insertions, 597 deletions
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 <host> XML tag representing a scanned host with
its services.
"""
- data = dict()
+ data = {}
# <status state="up" reason="user-set" reason_ttl="0"/>
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,