From ba544e2ba31432aad4d2acaba3e1f90c27137ded Mon Sep 17 00:00:00 2001 From: Max Nanis Date: Fri, 28 Aug 2026 00:22:17 -0700 Subject: Ruff evening! --- generalresearch/managers/events.py | 40 +++---- generalresearch/managers/gr/authentication.py | 3 +- generalresearch/managers/innovate/survey.py | 2 +- generalresearch/managers/leaderboard/manager.py | 8 +- generalresearch/managers/morning/survey.py | 2 +- generalresearch/managers/network/label.py | 2 +- generalresearch/managers/network/mtr.py | 9 +- generalresearch/managers/network/nmap.py | 9 +- generalresearch/managers/network/rdns.py | 5 +- generalresearch/managers/sago/survey.py | 2 +- generalresearch/managers/spectrum/survey.py | 2 +- generalresearch/managers/thl/buyer.py | 4 +- generalresearch/managers/thl/cashout_method.py | 4 +- generalresearch/managers/thl/category.py | 7 +- generalresearch/managers/thl/contest_manager.py | 9 +- generalresearch/managers/thl/ipinfo.py | 8 +- .../managers/thl/ledger_manager/ledger.py | 13 ++- generalresearch/managers/thl/payout.py | 46 ++++---- generalresearch/managers/thl/product.py | 33 +++--- generalresearch/managers/thl/profiling/user_upk.py | 8 +- generalresearch/managers/thl/survey.py | 121 ++++++++++----------- generalresearch/managers/thl/survey_penalty.py | 1 - generalresearch/managers/thl/tango_api.py | 2 +- .../managers/thl/user_manager/__init__.py | 2 +- .../thl/user_manager/mysql_user_manager.py | 11 +- generalresearch/managers/thl/userhealth.py | 8 +- generalresearch/managers/thl/wall.py | 2 +- generalresearch/managers/thl/wallet/tango.py | 6 +- generalresearch/mariadb.py | 8 -- generalresearch/models/admin/__init__.py | 2 +- generalresearch/models/admin/request.py | 2 +- generalresearch/models/cint/question.py | 2 +- generalresearch/models/cint/survey.py | 15 +-- generalresearch/models/custom_types.py | 6 +- generalresearch/models/dynata/survey.py | 12 +- generalresearch/models/dynata/task_collection.py | 2 +- generalresearch/models/gr/authentication.py | 13 +-- generalresearch/models/gr/business.py | 64 +++++------ generalresearch/models/gr/team.py | 19 ++-- generalresearch/models/innovate/question.py | 10 +- generalresearch/models/innovate/survey.py | 39 ++++--- generalresearch/models/legacy/bucket.py | 30 +++-- generalresearch/models/legacy/questions.py | 32 ++---- generalresearch/models/morning/survey.py | 10 +- generalresearch/models/morning/task_collection.py | 2 +- generalresearch/models/network/label.py | 12 +- generalresearch/models/network/nmap/result.py | 2 +- generalresearch/models/network/rdns/command.py | 2 +- generalresearch/models/precision/question.py | 6 +- generalresearch/models/precision/survey.py | 2 +- generalresearch/models/prodege/question.py | 11 +- generalresearch/models/prodege/survey.py | 6 +- generalresearch/models/prodege/task_collection.py | 2 +- generalresearch/models/repdata/question.py | 4 +- generalresearch/models/repdata/survey.py | 7 +- generalresearch/models/repdata/task_collection.py | 2 +- generalresearch/models/sago/question.py | 5 +- generalresearch/models/sago/survey.py | 24 ++-- generalresearch/models/thl/contest/__init__.py | 2 +- generalresearch/models/thl/contest/contest.py | 2 +- .../models/thl/contest/contest_entry.py | 1 + generalresearch/models/thl/contest/raffle.py | 6 +- generalresearch/models/thl/demographics.py | 8 +- generalresearch/models/thl/finance.py | 2 +- generalresearch/models/thl/ledger.py | 4 +- generalresearch/models/thl/offerwall/__init__.py | 2 +- generalresearch/models/thl/offerwall/base.py | 6 +- generalresearch/models/thl/payout_format.py | 2 +- generalresearch/models/thl/product.py | 4 +- .../models/thl/profiling/marketplace.py | 7 +- generalresearch/models/thl/report_task.py | 2 +- generalresearch/models/thl/session.py | 11 +- generalresearch/models/thl/soft_pair.py | 2 +- generalresearch/models/thl/survey/penalty.py | 2 +- .../models/thl/survey/task_collection.py | 3 +- generalresearch/models/thl/task_status.py | 9 +- generalresearch/pg_helper.py | 6 +- generalresearch/sql_helper.py | 2 +- generalresearch/utils/grpc_logger.py | 6 +- generalresearch/wall_status_codes/fullcircle.py | 2 +- generalresearch/wall_status_codes/innovate.py | 2 +- generalresearch/wall_status_codes/lucid.py | 2 +- generalresearch/wall_status_codes/morning.py | 2 +- generalresearch/wall_status_codes/pollfish.py | 2 +- test_utils/managers/contest/conftest.py | 4 +- test_utils/models/contest/conftest.py | 37 ++++--- test_utils/spectrum/conftest.py | 77 ++++++++++--- .../incite/collections/test_df_collection_base.py | 17 ++- .../collections/test_df_collection_item_base.py | 19 +++- .../test_df_collection_thl_marketplaces.py | 14 ++- .../collections/test_df_collection_thl_web.py | 120 ++++++++++++++------ .../mergers/foundations/test_user_id_product.py | 21 +++- tests/incite/mergers/test_pop_ledger.py | 6 +- tests/incite/test_collection_base.py | 2 +- tests/managers/thl/test_contest/test_milestone.py | 11 +- .../test_ledger/test_thl_lm_tx__user_payouts.py | 22 ++-- tests/models/spectrum/test_survey.py | 44 +------- 97 files changed, 669 insertions(+), 546 deletions(-) diff --git a/generalresearch/managers/events.py b/generalresearch/managers/events.py index f6c429e..efc8c0d 100644 --- a/generalresearch/managers/events.py +++ b/generalresearch/managers/events.py @@ -343,8 +343,8 @@ class TaskStatsManager(RedisManager): by_source=live_tasks_max_payout_by_source, ) - task_created_count_last_1h = dict() - task_created_count_last_24h = dict() + task_created_count_last_1h = {} + task_created_count_last_24h = {} for source in sources: task_created_count_last_1h[source] = pipe_res.pop(0) task_created_count_last_24h[source] = pipe_res.pop(0) @@ -381,25 +381,25 @@ class SessionStatsManager(RedisManager): older than 1 hr (in the 1 hr bucket) will expire. """ - # Must be ordered. Don't change this - global_keys = [ - "session_enters_last_1h", - "session_enters_last_24h", - "session_fails_last_1h", - "session_fails_last_24h", - "session_completes_last_1h", - "session_completes_last_24h", - "sum_payouts_last_1h", - "sum_payouts_last_24h", - "sum_user_payouts_last_1h", - "sum_user_payouts_last_24h", - # "session_fail_loi_sum_last_1h", - "session_fail_loi_sum_last_24h", - # "session_complete_loi_sum_last_1h", - "session_complete_loi_sum_last_24h", - ] - def __init__(self, *args, **kwargs): + # Must be ordered. Don't change this + self.global_keys = [ + "session_enters_last_1h", + "session_enters_last_24h", + "session_fails_last_1h", + "session_fails_last_24h", + "session_completes_last_1h", + "session_completes_last_24h", + "sum_payouts_last_1h", + "sum_payouts_last_24h", + "sum_user_payouts_last_1h", + "sum_user_payouts_last_24h", + # "session_fail_loi_sum_last_1h", + "session_fail_loi_sum_last_24h", + # "session_complete_loi_sum_last_1h", + "session_complete_loi_sum_last_24h", + ] + super().__init__(*args, **kwargs) self.SUM_HASH_LUA = self.redis_client.register_script(SUM_HASH_LUA_SCRIPT) diff --git a/generalresearch/managers/gr/authentication.py b/generalresearch/managers/gr/authentication.py index 80bee4b..721895e 100644 --- a/generalresearch/managers/gr/authentication.py +++ b/generalresearch/managers/gr/authentication.py @@ -270,7 +270,6 @@ class GRTokenManager(PostgresManager): ) conn.commit() - def get_by_user_id(self, user_id: PositiveInt) -> GRToken | None: # django authtoken_token table has (user_id) UNIQUE constraint # therefore, this will only return 0 or 1 GRTokens @@ -295,7 +294,7 @@ class GRTokenManager(PostgresManager): res = result[0] - for k, _ in res.items(): + for k in res: if isinstance(res[k], datetime): res[k] = res[k].replace(tzinfo=UTC) diff --git a/generalresearch/managers/innovate/survey.py b/generalresearch/managers/innovate/survey.py index cddfba2..f6d00a8 100644 --- a/generalresearch/managers/innovate/survey.py +++ b/generalresearch/managers/innovate/survey.py @@ -179,5 +179,5 @@ class InnovateSurveyManager(SurveyManager): if e.args[0] == 1062: existing_sns.add(sn) else: - raise e + raise self.update([surveys[sn] for sn in existing_sns]) diff --git a/generalresearch/managers/leaderboard/manager.py b/generalresearch/managers/leaderboard/manager.py index 07e3e2c..ed13cf2 100644 --- a/generalresearch/managers/leaderboard/manager.py +++ b/generalresearch/managers/leaderboard/manager.py @@ -45,9 +45,7 @@ class LeaderboardManager: self.country_iso = country_iso self.within_time_aware = None if within_time is None: - self.within_time_aware = datetime.now(tz=UTC).astimezone( - self.timezone - ) + self.within_time_aware = datetime.now(tz=UTC).astimezone(self.timezone) elif within_time.tzinfo is not None: self.within_time_aware = within_time.astimezone(self.timezone) else: @@ -123,7 +121,9 @@ class LeaderboardManager: user_idx = user_indices[0][0] user_row = user_indices[0][1] if user_row.rank == max([row.rank for row in rows]): - user_idx = [i for i, row in enumerate(rows) if row.rank == user_row.rank][0] + user_idx = next( + i for i, row in enumerate(rows) if row.rank == user_row.rank + ) start: int = max(user_idx - limit, 0) end: int = min(user_idx + limit + 1, len(rows)) diff --git a/generalresearch/managers/morning/survey.py b/generalresearch/managers/morning/survey.py index 0e29010..2488cd8 100644 --- a/generalresearch/managers/morning/survey.py +++ b/generalresearch/managers/morning/survey.py @@ -258,5 +258,5 @@ class MorningSurveyManager(SurveyManager): if e.args[0] == 1062: existing_sns.add(sn) else: - raise e + raise self.update([surveys[sn] for sn in existing_sns]) diff --git a/generalresearch/managers/network/label.py b/generalresearch/managers/network/label.py index 1f44862..f0ba9f7 100644 --- a/generalresearch/managers/network/label.py +++ b/generalresearch/managers/network/label.py @@ -30,7 +30,7 @@ class IPLabelManager(PostgresManager): params = ip_label.model_dump_postgres() with self.pg_config.make_connection() as conn, conn.cursor() as c: c.execute(query, params) - pk = c.fetchone()["id"] + _pk = c.fetchone()["id"] return ip_label def make_filter_str( diff --git a/generalresearch/managers/network/mtr.py b/generalresearch/managers/network/mtr.py index 19c5caf..54d74b7 100644 --- a/generalresearch/managers/network/mtr.py +++ b/generalresearch/managers/network/mtr.py @@ -42,8 +42,7 @@ class MTRRunManager(PostgresManager): if params_hops: c.executemany(query_hops, params_hops) else: - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute(query, params) - if params_hops: - c.executemany(query_hops, params_hops) + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute(query, params) + if params_hops: + c.executemany(query_hops, params_hops) diff --git a/generalresearch/managers/network/nmap.py b/generalresearch/managers/network/nmap.py index a8470c8..84d13ad 100644 --- a/generalresearch/managers/network/nmap.py +++ b/generalresearch/managers/network/nmap.py @@ -50,8 +50,7 @@ class NmapRunManager(PostgresManager): if nmap_run.ports: c.executemany(query_ports, params_ports) else: - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute(query, params) - if nmap_run.ports: - c.executemany(query_ports, params_ports) + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute(query, params) + if nmap_run.ports: + c.executemany(query_ports, params_ports) diff --git a/generalresearch/managers/network/rdns.py b/generalresearch/managers/network/rdns.py index 0b41a9a..c8ce913 100644 --- a/generalresearch/managers/network/rdns.py +++ b/generalresearch/managers/network/rdns.py @@ -28,6 +28,5 @@ class RDNSRunManager(PostgresManager): if c: c.execute(query, params) else: - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute(query, params) + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute(query, params) diff --git a/generalresearch/managers/sago/survey.py b/generalresearch/managers/sago/survey.py index 462d2ef..a13fbce 100644 --- a/generalresearch/managers/sago/survey.py +++ b/generalresearch/managers/sago/survey.py @@ -179,6 +179,6 @@ class SagoSurveyManager(SurveyManager): if e.args[0] == 1062: existing_sns.add(sn) else: - raise e + raise self.update([surveys[sn] for sn in existing_sns]) diff --git a/generalresearch/managers/spectrum/survey.py b/generalresearch/managers/spectrum/survey.py index 9b58d43..987ce7b 100644 --- a/generalresearch/managers/spectrum/survey.py +++ b/generalresearch/managers/spectrum/survey.py @@ -212,6 +212,6 @@ class SpectrumSurveyManager(SurveyManager): if e.args[0] == 1062: existing_sns.add(sn) else: - raise e + raise self.update([surveys[sn] for sn in existing_sns]) diff --git a/generalresearch/managers/thl/buyer.py b/generalresearch/managers/thl/buyer.py index 2cb582f..5aa2a01 100644 --- a/generalresearch/managers/thl/buyer.py +++ b/generalresearch/managers/thl/buyer.py @@ -18,8 +18,8 @@ class BuyerManager(PostgresManager): ): super().__init__(pg_config=pg_config, permissions=permissions) # self.buyer_pk: Dict[Buyer, int] = dict() - self.source_code_buyer: dict[str, Buyer] = dict() - self.source_code_pk: dict[str, int] = dict() + self.source_code_buyer: dict[str, Buyer] = {} + self.source_code_pk: dict[str, int] = {} self.populate_caches() def populate_caches(self): diff --git a/generalresearch/managers/thl/cashout_method.py b/generalresearch/managers/thl/cashout_method.py index 10282d4..f45e692 100644 --- a/generalresearch/managers/thl/cashout_method.py +++ b/generalresearch/managers/thl/cashout_method.py @@ -160,7 +160,7 @@ class CashoutMethodManager(PostgresManager): is_live: bool | None = True, ): filters = [] - params = dict() + params = {} if uuid is not None: params["uuid"] = uuid filters.append("id = %(uuid)s") @@ -292,7 +292,7 @@ class CashoutMethodManager(PostgresManager): x["type"] = PayoutType(x["provider"].upper()) if "data" not in x: - x["data"] = dict() + x["data"] = {} x["data"].update(x.pop("_data_")) x["data"]["type"] = x["type"] if user and x["type"] in {PayoutType.PAYPAL, PayoutType.CASH_IN_MAIL}: diff --git a/generalresearch/managers/thl/category.py b/generalresearch/managers/thl/category.py index 05ceb8f..e8a6aa6 100644 --- a/generalresearch/managers/thl/category.py +++ b/generalresearch/managers/thl/category.py @@ -9,8 +9,6 @@ from generalresearch.pg_helper import PostgresConfig class CategoryManager(PostgresManager): - categories = dict() - category_label_map = dict() def __init__( self, @@ -18,8 +16,9 @@ class CategoryManager(PostgresManager): permissions: Collection[Permission] | None = None, ): super().__init__(pg_config=pg_config, permissions=permissions) - self.categories: dict[UUIDStr, Category] = dict() - self.category_label_map: dict[str, Category] = dict() + self.categories: dict[UUIDStr, Category] = {} + self.category_label_map: dict[str, Category] = {} + self.populate_caches() def populate_caches(self): diff --git a/generalresearch/managers/thl/contest_manager.py b/generalresearch/managers/thl/contest_manager.py index 62146d7..64206e1 100644 --- a/generalresearch/managers/thl/contest_manager.py +++ b/generalresearch/managers/thl/contest_manager.py @@ -173,7 +173,7 @@ class ContestBaseManager(PostgresManager): except ValueError as e: if e.args[0] == "Contest not found": return None - raise e + raise @staticmethod def make_filter_str( @@ -187,7 +187,7 @@ class ContestBaseManager(PostgresManager): has_participants: bool | None = None, ) -> tuple[str, dict[str, Any]]: filters = [] - params = dict() + params = {} if product_id: params["product_id"] = product_id @@ -681,7 +681,7 @@ class RaffleContestManager(ContestBaseManager): raise ContestError(msg) if contest.entry_type == ContestEntryType.CASH: - tx = ledger_manager.create_tx_user_enter_contest( + ledger_manager.create_tx_user_enter_contest( contest_uuid=contest.uuid, contest_entry=entry ) @@ -827,7 +827,6 @@ class MilestoneContestManager(ContestBaseManager): ) self.end_milestone_contest(contest) - def enter_contest_db_work_milestone( self, contest: MilestoneUserView, user: User, incr: PositiveInt ) -> MilestoneEntry: @@ -1052,7 +1051,7 @@ class ContestManager( ) -> NonNegativeInt: contests_closed = 0 for contest in contests: - should_end, reason = contest.should_end() + should_end, _ = contest.should_end() if should_end: if hasattr(contest, "redis_client"): contest.redis_client = redis_client diff --git a/generalresearch/managers/thl/ipinfo.py b/generalresearch/managers/thl/ipinfo.py index e1143c2..9595cab 100644 --- a/generalresearch/managers/thl/ipinfo.py +++ b/generalresearch/managers/thl/ipinfo.py @@ -423,7 +423,9 @@ class IPInformationManager(PostgresManager): FROM thl_ipinformation WHERE updated >= NOW() - INTERVAL '12 hours' """ - denominator = list(pg_config.execute_sql_query(query=query))[0]["denominator"] + denominator = next(iter(pg_config.execute_sql_query(query=query)))[ + "denominator" + ] if denominator == 0: pass @@ -509,7 +511,7 @@ class GeoIpInfoManager(PostgresManagerWithRedis): res = [GeoIPInformation.model_validate_json(raw) for raw in res if raw] gs = {x.ip: x for x in res} - res2 = dict() + res2 = {} for ip, (normalized_ip, lookup_prefix) in ip_norm_lookup.items(): if normalized_ip not in gs: # try the non-normalized (remove me also 28 days from 2025-11-15) @@ -719,7 +721,7 @@ class GeoIpInfoManager(PostgresManagerWithRedis): gs = [GeoIPInformation.from_mysql(i) for i in res] gs = {g.ip: g for g in gs} - res2 = dict() + res2 = {} for ip, (normalized_ip, lookup_prefix) in ip_norm_lookup.items(): if normalized_ip not in gs: diff --git a/generalresearch/managers/thl/ledger_manager/ledger.py b/generalresearch/managers/thl/ledger_manager/ledger.py index 410f6ca..a5263f0 100644 --- a/generalresearch/managers/thl/ledger_manager/ledger.py +++ b/generalresearch/managers/thl/ledger_manager/ledger.py @@ -150,7 +150,7 @@ class LedgerTransactionManager(LedgerManagerBasePostgres): ), "LedgerTransactionManager has insufficient Permissions" if metadata is None: - metadata = dict() + metadata = {} if created is None: created = datetime.now(tz=UTC) @@ -429,7 +429,7 @@ class LedgerTransactionManager(LedgerManagerBasePostgres): ) } else: - metadata = dict() + metadata = {} entries = [ LedgerEntry( @@ -750,7 +750,7 @@ class LedgerMetadataManager(LedgerManagerBasePostgres): """ - tx_ids = set([tx.id for tx in transactions]) + tx_ids = {tx.id for tx in transactions} res = self.pg_config.execute_sql_query( query=""" SELECT @@ -782,7 +782,7 @@ class LedgerMetadataManager(LedgerManagerBasePostgres): from the database. """ - tx_ids = set([tx.id for tx in transactions]) + tx_ids = {tx.id for tx in transactions} res = self.pg_config.execute_sql_query( query=""" SELECT tx_meta.id @@ -792,7 +792,7 @@ class LedgerMetadataManager(LedgerManagerBasePostgres): params=[list(tx_ids)], ) - return set([i["id"] for i in res]) + return {i["id"] for i in res} class LedgerEntryManager(LedgerManagerBasePostgres): @@ -803,7 +803,7 @@ class LedgerEntryManager(LedgerManagerBasePostgres): def get_tx_entries_by_txs( self, transactions: list[LedgerTransaction] ) -> list[LedgerEntry]: - tx_ids = set([tx.id for tx in transactions]) + tx_ids = {tx.id for tx in transactions} tx_entries = self.pg_config.execute_sql_query( query=""" SELECT @@ -1141,4 +1141,5 @@ class LedgerManager( } for k, v in d.items(): v["total"] = (v["debit"] - v["credit"]) * k.normal_balance.value + return d diff --git a/generalresearch/managers/thl/payout.py b/generalresearch/managers/thl/payout.py index 8bc6843..f50e0d2 100644 --- a/generalresearch/managers/thl/payout.py +++ b/generalresearch/managers/thl/payout.py @@ -95,9 +95,9 @@ class PayoutEventManager(PostgresManagerWithRedis): with self.pg_config.make_connection() as conn: with conn.cursor() as c: c.execute(query=query, params=d) - assert c.rowcount == 1, ( - "Nothing was updated! Are you sure this payout_event exists?" - ) + assert ( + c.rowcount == 1 + ), "Nothing was updated! Are you sure this payout_event exists?" conn.commit() @@ -140,7 +140,7 @@ class UserPayoutEventManager(PayoutEventManager): # the purposes of returning to the user. pe = self.get_by_uuid(pe_uuid=pe_uuid) - transaction_info = dict() + transaction_info = {} order: dict[str, Any] = pe.order_data if pe.payout_type == PayoutType.TANGO and pe.status == PayoutStatus.COMPLETE: reward = order["reward"] @@ -411,7 +411,7 @@ class BrokerageProductPayoutEventManager(PayoutEventManager): *** IT IS ONLY FOR Brokerage Product PAYOUTS *** """ - params = dict() + params = {} filters = [] if ext_ref_id: # This is transaction id for tracking ACH/Wires with a banking institution @@ -630,9 +630,9 @@ class BusinessPayoutEventManager(PostgresManagerWithRedis): for bp_payout in d["bp_payouts"]: bp_payout["created"] = datetime.fromisoformat(bp_payout["created"]) bpe = BusinessPayoutEvent.model_validate(d) - assert bpe.bp_payouts is not None and len(bpe.bp_payouts) > 0, ( - "No BP payouts found for this Business Payout Event. This shouldn't happen!" - ) + assert ( + bpe.bp_payouts is not None and len(bpe.bp_payouts) > 0 + ), "No BP payouts found for this Business Payout Event. This shouldn't happen!" return bpe def filter_by( @@ -677,9 +677,9 @@ class BusinessPayoutEventManager(PostgresManagerWithRedis): for bp_payout in row["bp_payouts"]: bp_payout["created"] = datetime.fromisoformat(bp_payout["created"]) bpe = BusinessPayoutEvent.model_validate(row) - assert bpe.bp_payouts is not None and len(bpe.bp_payouts) > 0, ( - "No BP payouts found for this Business Payout Event. This shouldn't happen!" - ) + assert ( + bpe.bp_payouts is not None and len(bpe.bp_payouts) > 0 + ), "No BP payouts found for this Business Payout Event. This shouldn't happen!" bpes.append(bpe) return bpes @@ -696,9 +696,9 @@ class BusinessPayoutEventManager(PostgresManagerWithRedis): for bp_pe in bpe.bp_payouts ] txs = thl_lm.get_tx_ids_by_tags(tags=tags) - assert len(txs) == len(bpe.bp_payouts), ( - f"Expected {len(bpe.bp_payouts)} BP payouts but found {len(txs)}!" - ) + assert len(txs) == len( + bpe.bp_payouts + ), f"Expected {len(bpe.bp_payouts)} BP payouts but found {len(txs)}!" return True def resume_failed_business_payout( @@ -824,9 +824,9 @@ class BusinessPayoutEventManager(PostgresManagerWithRedis): shortfall: int = int(target_amount) - w_df["deduction"].sum() w_df["remaining_balance"] = w_df["available_balance"] - w_df["deduction"] - assert w_df[w_df["deduction"] > w_df["available_balance"]].empty, ( - "Trying to deduct more from an Product than what is available" - ) + assert w_df[ + w_df["deduction"] > w_df["available_balance"] + ].empty, "Trying to deduct more from an Product than what is available" return w_df @@ -898,9 +898,9 @@ class BusinessPayoutEventManager(PostgresManagerWithRedis): ) -> BusinessPayoutEvent: assert isinstance(bpe, BusinessPayoutEventCreate) assert bpe.bp_payouts, "Must provide at least one BP Payout" - assert {bp_pe.status for bp_pe in bpe.bp_payouts} == {PayoutStatus.PENDING}, ( - "All BP Payouts must be PENDING" - ) + assert {bp_pe.status for bp_pe in bpe.bp_payouts} == { + PayoutStatus.PENDING + }, "All BP Payouts must be PENDING" INSERT_SUPPLIER_PAYOUT = """ INSERT INTO supplier_payout ( business_id, created, amount, @@ -993,9 +993,9 @@ class BusinessPayoutEventManager(PostgresManagerWithRedis): # Can't pay any Products that don't have a remaining balance df = df[df["remaining_balance"] > 0].copy() - assert df.deduction.sum() == business.balance.recoup, ( - "recoup_proportional failure" - ) + assert ( + df.deduction.sum() == business.balance.recoup + ), "recoup_proportional failure" df["issue_amount"] = BusinessPayoutEventManager.distribute_amount( df=df, amount=amount diff --git a/generalresearch/managers/thl/product.py b/generalresearch/managers/thl/product.py index 54fa7c8..aac2979 100644 --- a/generalresearch/managers/thl/product.py +++ b/generalresearch/managers/thl/product.py @@ -33,6 +33,7 @@ if TYPE_CHECKING: ProfilingConfig, SessionConfig, SourcesConfig, + SupplyConfig, UserCreateConfig, UserHealthConfig, UserWalletConfig, @@ -167,15 +168,14 @@ class ProductManager(PostgresManager): if filter_uuids is None or len(filter_uuids) == 0: return [] - with self.pg_config.make_connection() as sql_connection: - with sql_connection.cursor() as c: - res = [] - for chunk in chunked(filter_uuids, 500): - res.extend( - self.fetch_uuids_( - c=c, filter_uuids=chunk, filter_column=filter_column - ) + with self.pg_config.make_connection() as sql_connection, sql_connection.cursor() as c: + res = [] + for chunk in chunked(filter_uuids, 500): + res.extend( + self.fetch_uuids_( + c=c, filter_uuids=chunk, filter_column=filter_column ) + ) return res def fetch_uuids_( @@ -258,9 +258,10 @@ class ProductManager(PostgresManager): for k, v in res1.items(): try: r.append(Product.model_validate(v)) - except ValidationError as e: + except ValidationError: logger.info(f"failed to parse product: {k}") - raise e + raise + return r def create( @@ -272,7 +273,7 @@ class ProductManager(PostgresManager): business_id: UUIDStr | None = None, harmonizer_domain: str | None = None, commission_pct: Decimal = Decimal("0.05"), - sources_config: SourcesConfig | SupplyConfigs | None = None, + sources_config: SourcesConfig | SupplyConfig | None = None, payout_config: PayoutConfig | None = None, session_config: SessionConfig | None = None, profiling_config: ProfilingConfig | None = None, @@ -360,10 +361,10 @@ class ProductManager(PostgresManager): insert_data["payments_enabled"] = instance.payments_enabled try: - insert_data["id_int"] = list(self.pg_config.execute_sql_query(query=""" + insert_data["id_int"] = next(iter(self.pg_config.execute_sql_query(query=""" SELECT COALESCE(MAX(id_int), 0) + 1 as id_int FROM userprofile_brokerageproduct - """))[0]["id_int"] + """)))["id_int"] instance.id_int = insert_data["id_int"] query = """ @@ -400,14 +401,14 @@ class ProductManager(PostgresManager): try: return self.get_by_uuid(product_uuid=instance.id) - except Exception: + except AssertionError: pass finally: self.cache_clear(instance.id) # If we couldn't find the Product, then go ahead and raise. capture_exception(e) - raise e + raise bpconfig = instance.model_dump( include={"sources_config", "user_wallet"}, mode="json" @@ -477,7 +478,7 @@ class ProductManager(PostgresManager): data["grs_domain"] = data.pop("harmonizer_domain") data = {k: v for k, v in data.items() if k in in_bp_keys} data["id"] = product_uuid - update_str = ", ".join(f"{k}=%({k})s" for k in data.keys()) + update_str = ", ".join(f"{k}=%({k})s" for k in data) self.pg_config.execute_write( f""" UPDATE userprofile_brokerageproduct diff --git a/generalresearch/managers/thl/profiling/user_upk.py b/generalresearch/managers/thl/profiling/user_upk.py index a2cddb3..6106037 100644 --- a/generalresearch/managers/thl/profiling/user_upk.py +++ b/generalresearch/managers/thl/profiling/user_upk.py @@ -158,7 +158,7 @@ class UserUpkManager(PostgresManagerWithRedis): country_isos = {x["country_iso"] for x in upk_ans_dict} assert len(country_isos) == 1 - country_iso = list(country_isos)[0] + country_iso = next(iter(country_isos)) for x in upk_ans_dict: x["pred"] = x["pred"].replace("gr:", "") x["obj"] = x["obj"].replace("gr:", "") @@ -304,15 +304,15 @@ class UserUpkManager(PostgresManagerWithRedis): def set_user_upk(self, upk_ans: list[UpkQuestionAnswer]) -> None: user_id = {x.user_id for x in upk_ans} assert len(user_id) == 1, "only run for 1 user at a time" - user_id = list(user_id)[0] + user_id = next(iter(user_id)) curr_upk = self.get_user_upk(user_id=user_id) curr_upk_simple = self.get_user_upk_simple(user_id=user_id) new_upk_simple = defaultdict(set) delete_items = set() - upk_multi = list() - delete_upk_multi = list() + upk_multi = [] + delete_upk_multi = [] for x in upk_ans: # For zero or more (multiple values) We want all values to equal these. # Might involve deleting values if they exist and are not in upk_ans diff --git a/generalresearch/managers/thl/survey.py b/generalresearch/managers/thl/survey.py index a9ec841..024ad38 100644 --- a/generalresearch/managers/thl/survey.py +++ b/generalresearch/managers/thl/survey.py @@ -134,7 +134,7 @@ class SurveyManager(PostgresManager): if len(survey_keys) == 0: return [] - params = dict() + params = {} survey_source_ids = defaultdict(set) for sk in survey_keys: @@ -354,59 +354,6 @@ class SurveyManager(PostgresManager): class SurveyStatManager(PostgresManager): - KEYS = [ - "survey_id", - "quota_id", - "country_iso", - "version", - "cpi", - "complete_too_fast_cutoff", - "prescreen_conv_alpha", - "prescreen_conv_beta", - "conv_alpha", - "conv_beta", - "dropoff_alpha", - "dropoff_beta", - "completion_time_mu", - "completion_time_sigma", - "mobile_eligible_alpha", - "mobile_eligible_beta", - "desktop_eligible_alpha", - "desktop_eligible_beta", - "tablet_eligible_alpha", - "tablet_eligible_beta", - "long_fail_rate", - "user_report_coeff", - "recon_likelihood", - "score_x0", - "score_x1", - "score", - "updated_at", - "survey_is_live", - "survey_survey_id", - "survey_source", - ] - - SURVEY_STATS_COL_MAP = { - "PRESCREEN_CONVERSION.alpha": "prescreen_conv_alpha", - "PRESCREEN_CONVERSION.beta": "prescreen_conv_beta", - "CONVERSION.alpha": "conv_alpha", - "CONVERSION.beta": "conv_beta", - "COMPLETION_TIME.mu": "completion_time_mu", - "COMPLETION_TIME.sigma": "completion_time_sigma", - "LONG_FAIL.value": "long_fail_rate", - "USER_REPORT_COEFF.value": "user_report_coeff", - "RECON_LIKELIHOOD.value": "recon_likelihood", - "DROPOFF_RATE.alpha": "dropoff_alpha", - "DROPOFF_RATE.beta": "dropoff_beta", - "IS_MOBILE_ELIGIBLE.alpha": "mobile_eligible_alpha", - "IS_MOBILE_ELIGIBLE.beta": "mobile_eligible_beta", - "IS_DESKTOP_ELIGIBLE.alpha": "desktop_eligible_alpha", - "IS_DESKTOP_ELIGIBLE.beta": "desktop_eligible_beta", - "IS_TABLET_ELIGIBLE.alpha": "tablet_eligible_alpha", - "IS_TABLET_ELIGIBLE.beta": "tablet_eligible_beta", - "cpi": "cpi", - } def __init__( self, @@ -419,6 +366,60 @@ class SurveyStatManager(PostgresManager): ) # self.ensure_surveystat_key_type() + self.KEYS = [ + "survey_id", + "quota_id", + "country_iso", + "version", + "cpi", + "complete_too_fast_cutoff", + "prescreen_conv_alpha", + "prescreen_conv_beta", + "conv_alpha", + "conv_beta", + "dropoff_alpha", + "dropoff_beta", + "completion_time_mu", + "completion_time_sigma", + "mobile_eligible_alpha", + "mobile_eligible_beta", + "desktop_eligible_alpha", + "desktop_eligible_beta", + "tablet_eligible_alpha", + "tablet_eligible_beta", + "long_fail_rate", + "user_report_coeff", + "recon_likelihood", + "score_x0", + "score_x1", + "score", + "updated_at", + "survey_is_live", + "survey_survey_id", + "survey_source", + ] + + self.SURVEY_STATS_COL_MAP = { + "PRESCREEN_CONVERSION.alpha": "prescreen_conv_alpha", + "PRESCREEN_CONVERSION.beta": "prescreen_conv_beta", + "CONVERSION.alpha": "conv_alpha", + "CONVERSION.beta": "conv_beta", + "COMPLETION_TIME.mu": "completion_time_mu", + "COMPLETION_TIME.sigma": "completion_time_sigma", + "LONG_FAIL.value": "long_fail_rate", + "USER_REPORT_COEFF.value": "user_report_coeff", + "RECON_LIKELIHOOD.value": "recon_likelihood", + "DROPOFF_RATE.alpha": "dropoff_alpha", + "DROPOFF_RATE.beta": "dropoff_beta", + "IS_MOBILE_ELIGIBLE.alpha": "mobile_eligible_alpha", + "IS_MOBILE_ELIGIBLE.beta": "mobile_eligible_beta", + "IS_DESKTOP_ELIGIBLE.alpha": "desktop_eligible_alpha", + "IS_DESKTOP_ELIGIBLE.beta": "desktop_eligible_beta", + "IS_TABLET_ELIGIBLE.alpha": "tablet_eligible_alpha", + "IS_TABLET_ELIGIBLE.beta": "tablet_eligible_beta", + "cpi": "cpi", + } + # # def ensure_surveystat_key_type(self): # SQL = """ @@ -570,12 +571,10 @@ class SurveyStatManager(PostgresManager): = (v.survey_id, v.quota_id, v.country_iso, v.version); """ params = [item for row in keys for item in row] - with self.pg_config.make_connection() as conn: - # self.register_surveystat_key(conn) - with conn.cursor() as c: - c.execute(query, params=params) - res = c.fetchall() - # print('\n'.join([x['QUERY PLAN'] for x in res])) + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute(query, params=params) + res = c.fetchall() + # print('\n'.join([x['QUERY PLAN'] for x in res])) return [SurveyStat.model_validate(x) for x in res] def update_surveystats_for_source( @@ -633,7 +632,7 @@ class SurveyStatManager(PostgresManager): country_iso: str | None = None, ) -> tuple[str, dict[str, Any]]: filters = [] - params = dict() + params = {} if updated_after is not None: params["updated_after"] = updated_after filters.append("ss.updated_at >= %(updated_after)s") diff --git a/generalresearch/managers/thl/survey_penalty.py b/generalresearch/managers/thl/survey_penalty.py index efaa930..bf914cb 100644 --- a/generalresearch/managers/thl/survey_penalty.py +++ b/generalresearch/managers/thl/survey_penalty.py @@ -61,7 +61,6 @@ class SurveyPenaltyManager(RedisManager): return f"{self.redis_prefix}:{uuid_id}" def set_penalties(self, penalties: list[Penalty]): - """ """ if len(penalties) > 1000: LOG.warning("SurveyPenaltyManager.set_penalties batch me!") assert len(penalties) < 10_000, "something is surely wrong" diff --git a/generalresearch/managers/thl/tango_api.py b/generalresearch/managers/thl/tango_api.py index 657224e..dab560e 100644 --- a/generalresearch/managers/thl/tango_api.py +++ b/generalresearch/managers/thl/tango_api.py @@ -122,7 +122,7 @@ class TangoClient: return self.get_order(reference_order_id) except TangoError as e: if "The order you requested cannot be found" not in e.args[0]: - raise e + raise return None def create_order(self, order: TangoOrderRequest) -> dict[str, Any]: diff --git a/generalresearch/managers/thl/user_manager/__init__.py b/generalresearch/managers/thl/user_manager/__init__.py index 0392edb..b3fa8f6 100644 --- a/generalresearch/managers/thl/user_manager/__init__.py +++ b/generalresearch/managers/thl/user_manager/__init__.py @@ -63,7 +63,7 @@ def parse_bp_trust_df(fp: str | Path) -> dict[str, Any]: "entrance_limit_value": convert_int, "median_daily_completes_7d": convert_int, } - bptrust = dict() + bptrust = {} with open(fp, newline="") as csvfile: reader = csv.reader(csvfile) diff --git a/generalresearch/managers/thl/user_manager/mysql_user_manager.py b/generalresearch/managers/thl/user_manager/mysql_user_manager.py index d2d0ffc..e0a7548 100644 --- a/generalresearch/managers/thl/user_manager/mysql_user_manager.py +++ b/generalresearch/managers/thl/user_manager/mysql_user_manager.py @@ -139,11 +139,10 @@ class MysqlUserManager: """) try: - with self.pg_config.make_connection() as conn: - with conn.cursor() as c: - c.execute(query=query, params=params) - user_id = c.fetchone()["id"] - except psycopg.IntegrityError as e: + with self.pg_config.make_connection() as conn, conn.cursor() as c: + c.execute(query=query, params=params) + user_id = c.fetchone()["id"] + except psycopg.IntegrityError: # Two machines/processes are trying to create this same (product_id, product_user_id) # at the same time. There's a unique index, so mysql will not let two be created. # The 2nd should get an IntegrityError, meaning this already exists, and we can just query it. @@ -160,7 +159,7 @@ class MysqlUserManager: else: # We specifically queried the NON read-replica, and we got an IntegrityError, so # something else must be wrong... - raise e + raise else: user = User( user_id=user_id, diff --git a/generalresearch/managers/thl/userhealth.py b/generalresearch/managers/thl/userhealth.py index babed04..26f08b4 100644 --- a/generalresearch/managers/thl/userhealth.py +++ b/generalresearch/managers/thl/userhealth.py @@ -221,7 +221,7 @@ class IPRecordManager(PostgresManagerWithRedis): "forwarded_ip5", "forwarded_ip6", ] - for col, ip in zip_longest( + for col, fwd_ip in zip_longest( fips_cols, [ forwarded_ip1, @@ -233,7 +233,7 @@ class IPRecordManager(PostgresManagerWithRedis): ], fillvalue=None, ): - data[col] = ipaddress.ip_address(ip).exploded if ip else ip + data[col] = ipaddress.ip_address(fwd_ip).exploded if fwd_ip else fwd_ip self.pg_config.execute_write( query=""" @@ -490,9 +490,7 @@ class AuditLogManager(PostgresManager): created_after: datetime | None = None, ) -> tuple[str, dict[str, Any]]: assert user_ids, "must pass at least 1 user_id" - assert all( - [isinstance(uid, int) for uid in user_ids] - ), "must pass user_id as int" + assert all(isinstance(uid, int) for uid in user_ids), "must pass user_id as int" if created_after is None: created_after = datetime.now(tz=UTC) - timedelta(days=7) diff --git a/generalresearch/managers/thl/wall.py b/generalresearch/managers/thl/wall.py index 03ca1c6..ac9fb62 100644 --- a/generalresearch/managers/thl/wall.py +++ b/generalresearch/managers/thl/wall.py @@ -484,7 +484,7 @@ class WallManager(PostgresManager): ORDER BY rs.source, rs.survey_id; """ - params = dict() + params = {} filters = [] # Instead of doing a big IN with a big set of tuples, since we know diff --git a/generalresearch/managers/thl/wallet/tango.py b/generalresearch/managers/thl/wallet/tango.py index 4abfc70..be8fd97 100644 --- a/generalresearch/managers/thl/wallet/tango.py +++ b/generalresearch/managers/thl/wallet/tango.py @@ -44,7 +44,7 @@ def complete_tango_order( tango_client=tango_client, ) - except Exception: + except AssertionError: # todo: its possible the order went through, but something else was wrong # we should try to retrieve the order by its ref_id and confirm it really # failed... @@ -70,8 +70,8 @@ def create_tango_order( """ Create a tango gift card order. Throws exception if anything is not right. - # https://integration-www.tangocard.com/raas_api_console/v2/ - # https://www.apimatic.io/apidocs/tangocard/v/2_3_4#/python + - https://integration-www.tangocard.com/raas_api_console/v2/ + - https://www.apimatic.io/apidocs/tangocard/v/2_3_4#/python :param utid: Card identifier :param amount: requested card value in USD diff --git a/generalresearch/mariadb.py b/generalresearch/mariadb.py index 8bcd8ee..5d43f97 100644 --- a/generalresearch/mariadb.py +++ b/generalresearch/mariadb.py @@ -32,11 +32,3 @@ def example(): for m in zip(c.metadata["field"], c.metadata["ext_type_or_format"]): # here we can just check if the field's ext_field_flag == 'UUID' (2) print(m[0], ext_field_flags_rev[m[1]]) - - -def get_column_types(): - # How does django do this? - res = """ - SELECT column_name, data_type - FROM information_schema.columns - WHERE table_name = 'morning_userpid' AND table_schema = DATABASE()""" diff --git a/generalresearch/models/admin/__init__.py b/generalresearch/models/admin/__init__.py index ad6302b..344c34a 100644 --- a/generalresearch/models/admin/__init__.py +++ b/generalresearch/models/admin/__init__.py @@ -1,6 +1,6 @@ from __future__ import annotations -from datetime import UTC, datetime, timezone +from datetime import UTC, datetime import pandas as pd from dateutil import relativedelta diff --git a/generalresearch/models/admin/request.py b/generalresearch/models/admin/request.py index 5fdc784..6112786 100644 --- a/generalresearch/models/admin/request.py +++ b/generalresearch/models/admin/request.py @@ -119,7 +119,7 @@ class ReportRequest(BaseModel): @property def end_naive(self) -> datetime: - return datetime.now(tz=None) + return datetime.now(tz=None) # noqa @property def ts_start(self) -> pd.Timestamp: diff --git a/generalresearch/models/cint/question.py b/generalresearch/models/cint/question.py index f8a287a..4c0e52f 100644 --- a/generalresearch/models/cint/question.py +++ b/generalresearch/models/cint/question.py @@ -44,7 +44,7 @@ class CintQuestionType(StrEnum): # This seems to be invalid as there are no options??? "Grid": None, } - return API_TYPE_MAP[a] if a in API_TYPE_MAP else None + return API_TYPE_MAP.get(a) class CintUserQuestionAnswer(MarketplaceUserQuestionAnswer): diff --git a/generalresearch/models/cint/survey.py b/generalresearch/models/cint/survey.py index cfc91ef..fde4559 100644 --- a/generalresearch/models/cint/survey.py +++ b/generalresearch/models/cint/survey.py @@ -12,6 +12,7 @@ from pydantic import ( ConfigDict, Field, NonNegativeInt, + ValidationError, computed_field, model_validator, ) @@ -67,7 +68,7 @@ class CintQuota(BaseModel): condition_hashes: list[str] | None = Field(min_length=1, default=None) def __hash__(self): - return hash(tuple((tuple(self.condition_hashes), self.quota_id))) + return hash((tuple(self.condition_hashes), self.quota_id)) @model_validator(mode="after") def validate_condition_len(self) -> Self: @@ -317,7 +318,7 @@ class CintSurvey(MarketplaceTask): def from_api(cls, d: dict[str, Any]) -> Self | 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 @@ -370,8 +371,8 @@ class CintSurvey(MarketplaceTask): d["mobile_conversion"] = None d["revenue_per_click"] = None - d["conditions"] = dict() - d.setdefault("survey_qualifications", list()) + d["conditions"] = {} + d.setdefault("survey_qualifications", []) qualifications = [CintCondition.from_api(q) for q in d["survey_qualifications"]] for q in qualifications: d["conditions"][q.criterion_hash] = q @@ -416,7 +417,7 @@ class CintSurvey(MarketplaceTask): return d @classmethod - def from_mysql(cls, d: Dict[str, Any]) -> Self: + def from_mysql(cls, d: dict[str, Any]) -> Self: d["created_at"] = d["created_at"].replace(tzinfo=UTC) d["last_updated"] = d["last_updated"].replace(tzinfo=UTC) d["qualifications"] = json.loads(d["qualifications"]) @@ -465,7 +466,7 @@ class CintSurvey(MarketplaceTask): ) -> tuple[bool | None, set[str]]: # Many surveys have 0 quotas. Quotas are exclusionary. # They can NOT match a quota where currently_open=0 - total_quota = [q for q in self.quotas if q.quota_type == "total"][0] + total_quota = next(q for q in self.quotas if q.quota_type == "total") if not total_quota.is_open: return False, set() quotas = [q for q in self.quotas if q.quota_type != "total"] @@ -474,7 +475,7 @@ class CintSurvey(MarketplaceTask): quota_eval = { quota: quota.matches_soft(criteria_evaluation) for quota in 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/custom_types.py b/generalresearch/models/custom_types.py index c200b34..5e4db3e 100644 --- a/generalresearch/models/custom_types.py +++ b/generalresearch/models/custom_types.py @@ -98,7 +98,7 @@ LanguageISOLike = Annotated[ def check_valid_uuid(v: str) -> str: try: assert UUID(v).hex == v - except Exception: + except (ValueError, AssertionError): raise ValueError("Invalid UUID") return v @@ -106,7 +106,7 @@ def check_valid_uuid(v: str) -> str: def is_valid_uuid(v: str) -> bool: try: assert UUID(v).hex == v - except Exception: + except (ValueError, AssertionError): return False return True @@ -165,7 +165,7 @@ CoercedStr = Annotated[str, BeforeValidator(coerce_int_to_str)] # Serializers that can transform a collection of str into a comma separated # str bidirectionally -to_comma_sep_str = PlainSerializer(lambda x: ",".join(sorted(list(x))), return_type=str) +to_comma_sep_str = PlainSerializer(lambda x: ",".join(sorted(x)), return_type=str) enum_to_comma_sep_str = PlainSerializer( lambda x: ",".join(sorted([str(y.value) for y in x])), return_type=str ) diff --git a/generalresearch/models/dynata/survey.py b/generalresearch/models/dynata/survey.py index 5a9f763..942ab4f 100644 --- a/generalresearch/models/dynata/survey.py +++ b/generalresearch/models/dynata/survey.py @@ -168,7 +168,7 @@ class DynataQuota(BaseModel): status: DynataStatus = Field() def __hash__(self): - return hash(tuple((tuple(self.condition_hashes), self.count, self.status))) + return hash((tuple(self.condition_hashes), self.count, self.status)) @property def is_open(self) -> bool: @@ -244,7 +244,7 @@ class DynataQuotaGroup(RootModel): ) -> tuple[bool | None, set[str]]: # Qualify for ANY quota object within a quota group obj_evals = {obj: obj.passes_soft(criteria_evaluation) for obj in self.root} - evals = set(v[0] for v in obj_evals.values()) + evals = {v[0] for v in obj_evals.values()} # If we match 1 obj, then the others don't matter if any(evals): return True, set() @@ -319,7 +319,7 @@ class DynataFilterGroup(RootModel): ) -> tuple[bool | None, set[str]]: # Passes back "passes" (T/F/none) and a list of unknown criterion hashes obj_evals = {obj: obj.passes_soft(criteria_evaluation) for obj in self.root} - evals = set(v[0] for v in obj_evals.values()) + evals = {v[0] for v in obj_evals.values()} # If we match 1 obj, then the others don't matter if any(evals): return True, set() @@ -548,7 +548,7 @@ class DynataSurvey(MarketplaceTask): return d @classmethod - def from_db(cls, d: Dict[str, Any]) -> Self: + def from_db(cls, d: dict[str, Any]) -> Self: d["created"] = d["created"].replace(tzinfo=UTC) d["last_updated"] = d["last_updated"].replace(tzinfo=UTC) d["filters"] = json.loads(d["filters"]) @@ -578,7 +578,7 @@ class DynataSurvey(MarketplaceTask): group_eval = { group: group.passes_soft(criteria_evaluation) for group in self.filters } - evals = set(g[0] for g in group_eval.values()) + evals = {g[0] for g in group_eval.values()} if False in evals: return False, set() elif None in evals: @@ -614,7 +614,7 @@ class DynataSurvey(MarketplaceTask): group_eval = { quota: quota.passes_soft(criteria_evaluation) for quota in self.quotas } - evals = set(g[0] for g in group_eval.values()) + evals = {g[0] for g in group_eval.values()} if False in evals: return False, set() elif None in evals: diff --git a/generalresearch/models/dynata/task_collection.py b/generalresearch/models/dynata/task_collection.py index 71cf3db..2b82bfd 100644 --- a/generalresearch/models/dynata/task_collection.py +++ b/generalresearch/models/dynata/task_collection.py @@ -54,7 +54,7 @@ DynataTaskCollectionSchema = DataFrameSchema( class DynataTaskCollection(TaskCollection): - items: List[DynataSurvey] + items: list[DynataSurvey] _schema = DynataTaskCollectionSchema def to_row(self, s: DynataSurvey) -> dict[str, Any]: diff --git a/generalresearch/models/gr/authentication.py b/generalresearch/models/gr/authentication.py index f9644fe..25f65fa 100644 --- a/generalresearch/models/gr/authentication.py +++ b/generalresearch/models/gr/authentication.py @@ -4,7 +4,7 @@ import binascii import json import os from datetime import UTC, datetime -from typing import TYPE_CHECKING, Any, Self +from typing import TYPE_CHECKING, Any from pydantic import ( AnyHttpUrl, @@ -283,16 +283,15 @@ class GRUser(BaseModel): ex=ex_secs, ) - # --- ORM --- @classmethod - def from_postgresql(cls, d: dict) -> Self: + def from_postgresql(cls, d: dict[str, Any]) -> GRUser: d["date_joined"] = d["date_joined"].replace(tzinfo=UTC) return GRUser.model_validate(d) @classmethod - def from_redis(cls, d: str | dict[str, Any]) -> Self: + def from_redis(cls, d: str | dict[str, Any]) -> GRUser: if isinstance(d, str): d = json.loads(d) assert isinstance(d, dict) @@ -357,13 +356,13 @@ class GRToken(BaseModel): # --- Properties --- @property - def auth_header(self, key_name: str = "Authorization") -> dict[str, str]: - return {key_name: self.key} + def auth_header(self) -> dict[str, str]: + return {"Authorization": self.key} # --- ORM --- @classmethod - def from_redis(cls, d: str | dict[str, Any]) -> Self: + def from_redis(cls, d: str | dict[str, Any]) -> GRToken: if isinstance(d, str): d = json.loads(d) assert isinstance(d, dict) diff --git a/generalresearch/models/gr/business.py b/generalresearch/models/gr/business.py index 74b5c29..064c200 100644 --- a/generalresearch/models/gr/business.py +++ b/generalresearch/models/gr/business.py @@ -6,14 +6,15 @@ import os from datetime import UTC, datetime from enum import Enum, StrEnum from pathlib import Path -from typing import TYPE_CHECKING, Self +from typing import TYPE_CHECKING from uuid import uuid4 import pandas as pd +import pyarrow as pa from dask.distributed import Client from psycopg.cursor import Cursor from psycopg.rows import dict_row -from pydantic import BaseModel, ConfigDict, Field, PositiveInt +from pydantic import BaseModel, ConfigDict, Field, PositiveInt, ValidationError from pydantic.json_schema import SkipJsonSchema from pydantic_extra_types.phone_numbers import PhoneNumber @@ -210,9 +211,11 @@ class Business(BaseModel): payouts: list[BusinessPayoutEvent] | None = Field( default=None, name="Business Payouts", - description="These are the ACH or Wire payments that were sent to the" - "Business as a single amount, summed for all the Business" - "child Products", + description=( + "These are the ACH or Wire payments that were sent to the" + "Business as a single amount, summed for all the Business" + "child Products" + ), ) pop_financial: list[POPFinancial] | None = Field(default=None) @@ -237,18 +240,19 @@ class Business(BaseModel): # --- Prefetch --- def prefetch_addresses(self, pg_config: PostgresConfig) -> None: - with pg_config.make_connection() as conn: - with conn.cursor(row_factory=dict_row) as c: - c.execute( - query=""" + with pg_config.make_connection() as conn, conn.cursor( + row_factory=dict_row + ) as c: + c.execute( + query=""" SELECT * FROM common_businessaddress AS ba WHERE ba.business_id = %s LIMIT 1 """, - params=[self.id], - ) - res = c.fetchall() + params=[self.id], + ) + res = c.fetchall() if len(res) == 0: self.addresses = [] @@ -258,22 +262,23 @@ class Business(BaseModel): def prefetch_teams(self, pg_config: PostgresConfig) -> None: from generalresearch.models.gr.team import Team - with pg_config.make_connection() as conn: - with conn.cursor(row_factory=dict_row) as c: - c: Cursor + with pg_config.make_connection() as conn, conn.cursor( + row_factory=dict_row + ) as c: + c: Cursor - c.execute( - query=""" + c.execute( + query=""" SELECT t.* FROM common_team AS t INNER JOIN common_team_businesses AS tb ON tb.team_id = t.id WHERE tb.business_id = %s """, - params=(self.id,), - ) + params=(self.id,), + ) - res = c.fetchall() + res = c.fetchall() if len(res) == 0: self.teams = [] @@ -542,11 +547,10 @@ class Business(BaseModel): ) try: - test = pd.read_parquet(path, engine="pyarrow") - except Exception as e: + _ = pd.read_parquet(path, engine="pyarrow") + except (pa.ArrowException, OSError, ValueError) as e: raise OSError(f"Parquet verification failed: {e}") - def prebuild_enriched_wall_parquet( self, thl_pg_config: PostgresConfig, @@ -586,11 +590,10 @@ class Business(BaseModel): ) try: - test = pd.read_parquet(path, engine="pyarrow") - except Exception as e: + _ = pd.read_parquet(path, engine="pyarrow") + except (pa.ArrowException, OSError, ValueError) as e: raise OSError(f"Parquet verification failed: {e}") - @classmethod def required_fields(cls) -> list[str]: return [ @@ -651,7 +654,7 @@ class Business(BaseModel): client=client, pop_ledger=pop_ledger, ) - self.prebuild_payouts(thl_pg_config=thl_web_rr, thl_lm=thl_lm, bpem=bpem) + self.prebuild_payouts(bpem=bpem) self.prebuild_pop_financial( thl_pg_config=thl_web_rr, thl_lm=thl_lm, @@ -713,7 +716,7 @@ class Business(BaseModel): uuid: UUIDStr, fields: list[str], gr_redis_config: RedisConfig, - ) -> Self | None: + ) -> Business | None: keys: list[str] = Business.required_fields() + fields if "pop_financial" in keys: @@ -724,7 +727,7 @@ class Business(BaseModel): rc = gr_redis_config.create_redis_client() try: - res: list = rc.hmget(name=f"business:{uuid}", keys=keys) + res: list[str | bytes | None] = rc.hmget(name=f"business:{uuid}", keys=keys) d = { val: json.loads(res[idx]) if res[idx] is not None else None for idx, val in enumerate(keys) @@ -742,6 +745,5 @@ class Business(BaseModel): result["pop_financial"] = pop_financial return Business.model_validate(result) - except Exception as e: - logging.exception(e) + except ValidationError: return None diff --git a/generalresearch/models/gr/team.py b/generalresearch/models/gr/team.py index 78a9ba9..4752bea 100644 --- a/generalresearch/models/gr/team.py +++ b/generalresearch/models/gr/team.py @@ -5,16 +5,18 @@ import os from datetime import UTC, datetime from enum import Enum from pathlib import Path -from typing import TYPE_CHECKING, Self +from typing import TYPE_CHECKING from uuid import uuid4 import pandas as pd +import pyarrow as pa from dask.distributed import Client from pydantic import ( BaseModel, ConfigDict, Field, PositiveInt, + ValidationError, field_validator, ) from pydantic.json_schema import SkipJsonSchema @@ -191,10 +193,9 @@ class Team(BaseModel): try: _ = pd.read_parquet(path, engine="pyarrow") - except Exception as e: + except (pa.ArrowException, OSError, ValueError) as e: raise OSError(f"Parquet verification failed: {e}") - def prebuild_enriched_wall_parquet( self, thl_pg_config: PostgresConfig, @@ -235,10 +236,9 @@ class Team(BaseModel): try: _ = pd.read_parquet(path, engine="pyarrow") - except Exception as e: + except (pa.ArrowException, OSError, ValueError) as e: raise OSError(f"Parquet verification failed: {e}") - @classmethod def required_fields(cls) -> list[str]: return [ @@ -281,8 +281,6 @@ class Team(BaseModel): enriched_session: EnrichedSessionMerge | None = None, enriched_wall: EnrichedWallMerge | None = None, ) -> None: - ex_secs = 60 * 60 * 24 * 3 # 3 days - self.prefetch_products(thl_pg_config=thl_web_rr) self.prefetch_gr_users(pg_config=pg_config, redis_config=redis_config) self.prefetch_businesses(pg_config=pg_config, redis_config=redis_config) @@ -323,7 +321,6 @@ class Team(BaseModel): enriched_wall=enriched_wall, ) - # --- ORM --- @classmethod @@ -332,14 +329,14 @@ class Team(BaseModel): uuid: UUIDStr, fields: list[str], gr_redis_config: RedisConfig, - ) -> Self | None: + ) -> Team | None: keys: list = Team.required_fields() + fields rc = gr_redis_config.create_redis_client() try: - res: list = rc.hmget(name=f"team:{uuid}", keys=keys) + res: list[str | bytes | None] = rc.hmget(name=f"team:{uuid}", keys=keys) d = {val: json.loads(res[idx]) for idx, val in enumerate(keys)} return Team.model_validate(d) - except Exception: + except ValidationError: return None diff --git a/generalresearch/models/innovate/question.py b/generalresearch/models/innovate/question.py index 89310a2..f5a4846 100644 --- a/generalresearch/models/innovate/question.py +++ b/generalresearch/models/innovate/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 from generalresearch.models.innovate import InnovateQuestionID @@ -71,7 +71,7 @@ class InnovateQuestionType(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) class InnovateQuestion(MarketplaceQuestion): @@ -141,7 +141,7 @@ class InnovateQuestion(MarketplaceQuestion): @classmethod def from_api( - cls, d: dict, country_iso: str, language_iso: str + cls, d: dict[str, Any], country_iso: str, language_iso: str ) -> InnovateQuestion | None: """ :param d: Raw response from API @@ -151,13 +151,13 @@ class InnovateQuestion(MarketplaceQuestion): """ 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 @classmethod def _from_api( - cls, d: dict, country_iso: str, language_iso: str + cls, d: dict[str, Any], country_iso: str, language_iso: str ) -> InnovateQuestion: # Question AGE returns options even though its marked as a text entry (but only in some locales) d["QuestionKey"] = d["QuestionKey"].lower() diff --git a/generalresearch/models/innovate/survey.py b/generalresearch/models/innovate/survey.py index 3c37fe3..d07f960 100644 --- a/generalresearch/models/innovate/survey.py +++ b/generalresearch/models/innovate/survey.py @@ -17,6 +17,7 @@ from pydantic import ( BaseModel, ConfigDict, Field, + ValidationError, computed_field, model_validator, ) @@ -69,7 +70,7 @@ class InnovateCondition(MarketplaceCondition): d["logical_operator"] = LogicalOperator.OR d["value_type"] = ConditionValueType.LIST d["negate"] = False - d["values"] = list(set(x.strip().lower() for x in d["values"])) + d["values"] = list({x.strip().lower() for x in d["values"]}) return cls.model_validate(d) @@ -88,7 +89,7 @@ class InnovateQuota(BaseModel): condition_hashes: list[str] = Field(min_length=0, default_factory=list) def __hash__(self): - 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: @@ -99,7 +100,7 @@ class InnovateQuota(BaseModel): ) @classmethod - def from_api(cls, d: dict): + def from_api(cls, d: dict[str, Any]): return cls.model_validate(d) def passes(self, criteria_evaluation: dict[str, bool | None]) -> bool: @@ -263,13 +264,13 @@ class InnovateSurvey(MarketplaceTask): def from_api(cls, d: dict[str, Any]) -> InnovateSurvey | 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 @classmethod def _from_api(cls, d: dict[str, Any]) -> InnovateSurvey: - d["conditions"] = dict() + d["conditions"] = {} # If we haven't hit the "detail" endpoint, we won't get this d.setdefault("qualifications", []) @@ -317,11 +318,14 @@ class InnovateSurvey(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 {"exclude_pids", "include_pids", "excluded_surveys"}: - if v and len(v) > 6: - v = sorted(v) - v = v[:3] + ["…"] + v[-3:] - repr_args[n] = (k, v) + if ( + k in {"exclude_pids", "include_pids", "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 @@ -380,14 +384,21 @@ class InnovateSurvey(MarketplaceTask): """ assert isinstance(att_survey_ids, set), "must pass a set" assert isinstance(att_job_ids, set), "must pass a set" + if self.survey_id in att_survey_ids: return False - if self.duplicate_check_level == InnovateDuplicateCheckLevel.JOB: - if self.job_id in att_job_ids: - return False + + if ( + self.duplicate_check_level == InnovateDuplicateCheckLevel.JOB + and self.job_id in att_job_ids + ): + return False + if self.duplicate_check_level == InnovateDuplicateCheckLevel.EXCLUDED_SURVEYS: + assert self.excluded_surveys is not None if self.excluded_surveys.intersection(att_survey_ids): return False + return True def passes_qualifications( @@ -431,7 +442,7 @@ class InnovateSurvey(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/legacy/bucket.py b/generalresearch/models/legacy/bucket.py index 812241d..3705f4f 100644 --- a/generalresearch/models/legacy/bucket.py +++ b/generalresearch/models/legacy/bucket.py @@ -120,8 +120,10 @@ class BucketBase(BaseModel): ) uri: HttpsUrl = Field( examples=[ - "https://task.generalresearch.com/api/v1/52d3f63b2709/797df4136c604a6c8599818296aae6d1/?i" - "=5ba2fe5010cc4d078fc3cc0b0cc264c3&b=test&66482fb=e7baf5e" + ( + "https://task.generalresearch.com/api/v1/52d3f63b2709/797df4136c604a6c8599818296aae6d1/?i" + "=5ba2fe5010cc4d078fc3cc0b0cc264c3&b=test&66482fb=e7baf5e" + ) ], description="The URL to send a respondent into. Must not edit this URL in any way", ) @@ -465,12 +467,12 @@ class PayoutSummaryDecimal(StatisticalSummary): class PayoutSummary(StatisticalSummary): """Payouts are in Integer USD Cents""" - min: int = Field(gt=0, le=10000) - max: int = Field(gt=0, le=10000) - q1: int = Field(gt=0, le=10000) - q2: int = Field(gt=0, le=10000) - q3: int = Field(gt=0, le=10000) - mean: int | None = Field(gt=0, le=10000, default=None) + min: int = Field(gt=0, le=10_000) + max: int = Field(gt=0, le=10_000) + q1: int = Field(gt=0, le=10_000) + q2: int = Field(gt=0, le=10_000) + q3: int = Field(gt=0, le=10_000) + mean: int | None = Field(gt=0, le=10_000, default=None) model_config = { "json_schema_extra": { @@ -724,8 +726,10 @@ class OneShotOfferwallBucket(BaseModel): ) uri: HttpsUrl = Field( examples=[ - "https://task.generalresearch.com/api/v1/52d3f63b2709/797df4136c604a6c8599818296aae6d1/?i" - "=5ba2fe5010cc4d078fc3cc0b0cc264c3&b=test&66482fb=e7baf5e" + ( + "https://task.generalresearch.com/api/v1/52d3f63b2709/797df4136c604a6c8599818296aae6d1/?i" + "=5ba2fe5010cc4d078fc3cc0b0cc264c3&b=test&66482fb=e7baf5e" + ) ], description="The URL to send a respondent into. Must not edit this URL in any way", ) @@ -759,8 +763,10 @@ class WXETOfferwallBucket(BaseModel): ) uri: HttpsUrl = Field( examples=[ - "https://task.generalresearch.com/api/v1/52d3f63b2709/797df4136c604a6c8599818296aae6d1/?i" - "=5ba2fe5010cc4d078fc3cc0b0cc264c3&b=test&66482fb=e7baf5e" + ( + "https://task.generalresearch.com/api/v1/52d3f63b2709/797df4136c604a6c8599818296aae6d1/?i" + "=5ba2fe5010cc4d078fc3cc0b0cc264c3&b=test&66482fb=e7baf5e" + ) ], description="The URL to send a respondent into. Must not edit this URL in any way", ) diff --git a/generalresearch/models/legacy/questions.py b/generalresearch/models/legacy/questions.py index 4651ab0..9f37837 100644 --- a/generalresearch/models/legacy/questions.py +++ b/generalresearch/models/legacy/questions.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Annotated, Any, Self +from typing import TYPE_CHECKING, Annotated, Any from pydantic import ( BaseModel, @@ -87,25 +87,17 @@ class UserQuestionAnswerIn(BaseModel): fingerprint_tz = "a91cb1dea814480dba12d9b7b48696dd" fingerprint_fingerprint = "1d1e2e8380ac474b87fb4e4c569b48df" - if self.question_id in { - user_agent_qid, - fingerprint_langs, - fingerprint_tz, - fingerprint_fingerprint, - }: - if len(self.answer) != 1: - raise ValueError("Too many answer values provided") - - return self - - @model_validator(mode="after") - def user_agent_check(self) -> Self: - # TODO: where / how do I want to pass in this Werz user_agent stuff? - user_agent_qid = "2fbedb2b9f7647b09ff5e52fa119cc5e" - - if self.question_id == user_agent_qid: - val = self.answer[0] - # assert val == request.user_agent.to_header(): + if ( + self.question_id + in { + user_agent_qid, + fingerprint_langs, + fingerprint_tz, + fingerprint_fingerprint, + } + and len(self.answer) != 1 + ): + raise ValueError("Too many answer values provided") return self diff --git a/generalresearch/models/morning/survey.py b/generalresearch/models/morning/survey.py index 5011698..91d1bce 100644 --- a/generalresearch/models/morning/survey.py +++ b/generalresearch/models/morning/survey.py @@ -200,7 +200,7 @@ class MorningQuota(MorningStatistics, MarketplaceTask): data["country_isos"] = [data["country_iso"]] if isinstance(data["language_isos"], str): data["language_isos"] = set(data["language_isos"].split(",")) - data["language_iso"] = sorted(data["language_isos"])[0] + data["language_iso"] = min(data["language_isos"]) return data @property @@ -276,11 +276,11 @@ class MorningQuota(MorningStatistics, MarketplaceTask): 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: @@ -359,7 +359,7 @@ class MorningBid(MorningTaskStatistics): @property def language_iso_any(self): - return sorted(self.language_isos)[0] + return min(self.language_isos) @property def locale(self): @@ -417,7 +417,7 @@ class MorningBid(MorningTaskStatistics): if "conditions" in data: return data - data["conditions"] = dict() + data["conditions"] = {} for quota in data["quotas"]: if "qualifications" in quota: quota_conditions = [ diff --git a/generalresearch/models/morning/task_collection.py b/generalresearch/models/morning/task_collection.py index eb4cbd1..9303a2f 100644 --- a/generalresearch/models/morning/task_collection.py +++ b/generalresearch/models/morning/task_collection.py @@ -108,7 +108,7 @@ class MorningTaskCollection(TaskCollection): ] quota_fields = list(quota_columns.keys()) rows = [] - bid_dict = dict() + bid_dict = {} for k in bid_fields: bid_dict[k] = getattr(bid, k) bid_dict["bid.id"] = bid.id diff --git a/generalresearch/models/network/label.py b/generalresearch/models/network/label.py index e4ddd18..60a6e58 100644 --- a/generalresearch/models/network/label.py +++ b/generalresearch/models/network/label.py @@ -2,6 +2,7 @@ from __future__ import annotations import ipaddress from enum import StrEnum +from ipaddress import IPv4Network, IPv6Network from pydantic import ( BaseModel, @@ -84,12 +85,13 @@ class IPLabel(BaseModel): @field_validator("ip", mode="before") @classmethod - def normalize_and_validate_network(cls, v): - net = ipaddress.ip_network(v, strict=False) + def normalize_and_validate_network( + cls, v: IPvAnyNetwork + ) -> IPv4Network | IPv6Network | None: + net = ipaddress.ip_network(address=v, strict=False) - if isinstance(net, ipaddress.IPv6Network): - if net.prefixlen > 64: - raise ValueError("IPv6 network must be /64 or larger") + if isinstance(net, ipaddress.IPv6Network) and net.prefixlen > 64: + raise ValueError("IPv6 network must be /64 or larger") return net diff --git a/generalresearch/models/network/nmap/result.py b/generalresearch/models/network/nmap/result.py index 55c2109..4552e15 100644 --- a/generalresearch/models/network/nmap/result.py +++ b/generalresearch/models/network/nmap/result.py @@ -411,7 +411,7 @@ class NmapResult(BaseModel): def model_dump_postgres(self): # Writes for the network_portscan table - d = dict() + d = {} data = self.model_dump( mode="json", include={ diff --git a/generalresearch/models/network/rdns/command.py b/generalresearch/models/network/rdns/command.py index e88a84d..bccead0 100644 --- a/generalresearch/models/network/rdns/command.py +++ b/generalresearch/models/network/rdns/command.py @@ -20,7 +20,7 @@ def run_rdns(config: RDNSRunCommand) -> RDNSResult: def build_rdns_command(ip: str) -> str: # e.g. dig +noall +answer -x 1.2.3.4 - return " ".join(["dig", "+noall", "+answer", "-x", ip]) + return f"dig +noall +answer -x {ip}" def get_dig_version() -> str: diff --git a/generalresearch/models/precision/question.py b/generalresearch/models/precision/question.py index cc90aa9..f532998 100644 --- a/generalresearch/models/precision/question.py +++ b/generalresearch/models/precision/question.py @@ -54,15 +54,15 @@ class PrecisionQuestionType(StrEnum): TEXT_ENTRY = "t" @classmethod - def from_api(cls, a: int): - API_TYPE_MAP = { + def from_api(cls, a: int) -> PrecisionQuestionType | None: + api_type_map: dict[str, PrecisionQuestionType] = { "Drop Down": PrecisionQuestionType.SINGLE_SELECT, "Multi Select": PrecisionQuestionType.MULTI_SELECT, "Single Select": PrecisionQuestionType.SINGLE_SELECT, "Single Select Matrix": PrecisionQuestionType.SINGLE_SELECT, "Vertical Question": PrecisionQuestionType.SINGLE_SELECT, } - return API_TYPE_MAP[a] if a in API_TYPE_MAP else None + return api_type_map.get(a, None) class PrecisionUserQuestionAnswer(MarketplaceUserQuestionAnswer): diff --git a/generalresearch/models/precision/survey.py b/generalresearch/models/precision/survey.py index b27b8c4..bf9e83e 100644 --- a/generalresearch/models/precision/survey.py +++ b/generalresearch/models/precision/survey.py @@ -312,7 +312,7 @@ class PrecisionSurvey(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/prodege/question.py b/generalresearch/models/prodege/question.py index 574c4fd..58aed67 100644 --- a/generalresearch/models/prodege/question.py +++ b/generalresearch/models/prodege/question.py @@ -8,7 +8,14 @@ from enum import StrEnum from functools import cached_property from typing import TYPE_CHECKING, Any, Literal -from pydantic import BaseModel, ConfigDict, Field, PositiveInt, model_validator +from pydantic import ( + BaseModel, + ConfigDict, + Field, + PositiveInt, + ValidationError, + model_validator, +) from generalresearch.locales import Localelator from generalresearch.models import MAX_INT32, Source @@ -143,7 +150,7 @@ class ProdegeQuestion(MarketplaceQuestion): """ try: return cls._from_api(d, country_iso) - except Exception as e: + except ValidationError as e: logger.warning(f"Unable to parse question: {d}. {e}") return None diff --git a/generalresearch/models/prodege/survey.py b/generalresearch/models/prodege/survey.py index 7ab6df6..5d0369a 100644 --- a/generalresearch/models/prodege/survey.py +++ b/generalresearch/models/prodege/survey.py @@ -539,7 +539,7 @@ class ProdegeSurvey(MarketplaceTask): d["country_isos"] = [ locale_helper.get_country_iso(d.pop("country_code").lower()) ] - d["country_iso"] = sorted(d["country_isos"])[0] + d["country_iso"] = min(d["country_isos"]) # No languages are returned anywhere for anything d["language_isos"] = [ locale_helper.get_default_lang_from_country(d["country_isos"][0]) @@ -552,7 +552,7 @@ class ProdegeSurvey(MarketplaceTask): d["past_participation"] = ProdegePastParticipation.from_api( d["past_participation"] ) - d["conditions"] = dict() + d["conditions"] = {} for quota in d["quotas"]: quota["condition_hashes"] = [] for c in quota["targeting_criteria"]: @@ -563,7 +563,7 @@ class ProdegeSurvey(MarketplaceTask): d["quotas"] = [ProdegeQuota.from_api(q) for q in d["quotas"]] countries = {q.country_iso for q in d["quotas"] if q.country_iso} if countries: - d["country_iso"] = sorted(countries)[0] + d["country_iso"] = min(countries) d["country_isos"] = countries d["language_iso"] = locale_helper.get_default_lang_from_country( d["country_iso"] diff --git a/generalresearch/models/prodege/task_collection.py b/generalresearch/models/prodege/task_collection.py index 4544050..9f6a81b 100644 --- a/generalresearch/models/prodege/task_collection.py +++ b/generalresearch/models/prodege/task_collection.py @@ -76,7 +76,7 @@ class ProdegeTaskCollection(TaskCollection): "used_question_ids", "all_hashes", ] - d = dict() + d = {} for k in fields: d[k] = getattr(s, k) d["cpi"] = float(d["cpi"]) diff --git a/generalresearch/models/repdata/question.py b/generalresearch/models/repdata/question.py index 9dda97f..4fa2d22 100644 --- a/generalresearch/models/repdata/question.py +++ b/generalresearch/models/repdata/question.py @@ -12,6 +12,7 @@ from pydantic import ( ConfigDict, Field, PositiveInt, + ValidationError, field_validator, model_validator, ) @@ -142,6 +143,7 @@ class RepDataQuestion(MarketplaceQuestion): @property def internal_id(self) -> str: + assert self.lucid_id return self.lucid_id @field_validator("question_id", mode="before") @@ -167,7 +169,7 @@ class RepDataQuestion(MarketplaceQuestion): """ 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 diff --git a/generalresearch/models/repdata/survey.py b/generalresearch/models/repdata/survey.py index 43a592c..c5b0730 100644 --- a/generalresearch/models/repdata/survey.py +++ b/generalresearch/models/repdata/survey.py @@ -13,6 +13,7 @@ from pydantic import ( BaseModel, ConfigDict, Field, + ValidationError, computed_field, field_validator, model_validator, @@ -459,7 +460,7 @@ class RepDataSurvey(BaseModel): @property def all_conditions(self) -> list[RepDataCondition]: - cs = list() + cs = [] for stream in self.streams: cs.extend(stream.all_conditions) # dedupe by criterion_hash @@ -477,7 +478,7 @@ class RepDataSurvey(BaseModel): """ try: return cls._from_api(survey_response) - except Exception as e: + except ValidationError as e: survey_id = survey_response.get("survey_id") or survey_response.get( "SurveyNumber" ) @@ -485,7 +486,7 @@ class RepDataSurvey(BaseModel): return None @classmethod - def _from_api(cls, survey_response) -> RepDataSurvey: + def _from_api(cls, survey_response: dict[str, Any]) -> RepDataSurvey: d = survey_response.copy() d["country_iso"] = locale_helper.get_country_iso(d["SurveyCountry"].lower()) d["language_iso"] = locale_helper.get_language_iso(d["SurveyLanguage"].lower()) diff --git a/generalresearch/models/repdata/task_collection.py b/generalresearch/models/repdata/task_collection.py index d625349..5b9a4ba 100644 --- a/generalresearch/models/repdata/task_collection.py +++ b/generalresearch/models/repdata/task_collection.py @@ -110,7 +110,7 @@ class RepDataTaskCollection(TaskCollection): "remaining_count", ] rows = [] - d = dict() + d = {} for k in survey_fields: d[k] = getattr(s, k) d["allowed_devices"] = s.allowed_devices_str diff --git a/generalresearch/models/sago/question.py b/generalresearch/models/sago/question.py index 474543d..291214f 100644 --- a/generalresearch/models/sago/question.py +++ b/generalresearch/models/sago/question.py @@ -13,6 +13,7 @@ from pydantic import ( ConfigDict, Field, PositiveInt, + ValidationError, field_validator, model_validator, ) @@ -86,7 +87,7 @@ class SagoQuestionType(StrEnum): 6: SagoQuestionType.TEXT_ENTRY, 7: SagoQuestionType.TEXT_ENTRY, } - return API_TYPE_MAP[a] if a in API_TYPE_MAP else None + return API_TYPE_MAP.get(a, None) class SagoUserQuestionAnswer(BaseModel): @@ -182,7 +183,7 @@ class SagoQuestion(MarketplaceQuestion): """ 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 diff --git a/generalresearch/models/sago/survey.py b/generalresearch/models/sago/survey.py index 83aad8c..8550cd3 100644 --- a/generalresearch/models/sago/survey.py +++ b/generalresearch/models/sago/survey.py @@ -8,7 +8,14 @@ from functools import cached_property from typing import Annotated, Any, Literal, Self from more_itertools import flatten -from pydantic import BaseModel, ConfigDict, Field, computed_field, model_validator +from pydantic import ( + BaseModel, + ConfigDict, + Field, + ValidationError, + computed_field, + model_validator, +) from generalresearch.locales import Localelator from generalresearch.models import LogicalOperator, Source @@ -71,7 +78,7 @@ class SagoQuota(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: @@ -261,7 +268,7 @@ class SagoSurvey(MarketplaceTask): def from_api(cls, d: dict[str, Any]) -> SagoSurvey | 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 @@ -273,11 +280,10 @@ class SagoSurvey(MarketplaceTask): # Fancy repr that abbreviates ip_exclusions and survey_exclusions repr_args = list(self.__repr_args__()) for n, (k, v) in enumerate(repr_args): - if k in {"ip_exclusions", "survey_exclusions"}: - if v and len(v) > 6: - v = sorted(v) - v = v[:3] + ["…"] + v[-3:] - repr_args[n] = (k, v) + if k in {"ip_exclusions", "survey_exclusions"} 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 @@ -362,7 +368,7 @@ class SagoSurvey(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/thl/contest/__init__.py b/generalresearch/models/thl/contest/__init__.py index 842d85e..65b28f4 100644 --- a/generalresearch/models/thl/contest/__init__.py +++ b/generalresearch/models/thl/contest/__init__.py @@ -1,7 +1,7 @@ from __future__ import annotations from datetime import UTC, datetime -from typing import Self +from typing import Any, Self from uuid import uuid4 from pydantic import ( diff --git a/generalresearch/models/thl/contest/contest.py b/generalresearch/models/thl/contest/contest.py index bd0fc04..5814bef 100644 --- a/generalresearch/models/thl/contest/contest.py +++ b/generalresearch/models/thl/contest/contest.py @@ -186,7 +186,7 @@ class Contest(ContestBase): @classmethod def model_validate_mysql(cls, data: dict[str, Any]) -> Self: - data = {k: v for k, v in data.items() if k in cls.model_fields.keys()} + data = {k: v for k, v in data.items() if k in cls.model_fields} if isinstance(data["end_condition"], dict): data["end_condition"] = ContestEndCondition.model_validate( data["end_condition"] diff --git a/generalresearch/models/thl/contest/contest_entry.py b/generalresearch/models/thl/contest/contest_entry.py index bb3aef4..b5f0ac3 100644 --- a/generalresearch/models/thl/contest/contest_entry.py +++ b/generalresearch/models/thl/contest/contest_entry.py @@ -1,6 +1,7 @@ from __future__ import annotations from datetime import UTC, datetime +from typing import Any from uuid import uuid4 from pydantic import ( diff --git a/generalresearch/models/thl/contest/raffle.py b/generalresearch/models/thl/contest/raffle.py index 16a0a47..072f011 100644 --- a/generalresearch/models/thl/contest/raffle.py +++ b/generalresearch/models/thl/contest/raffle.py @@ -203,9 +203,7 @@ class RaffleContest(RaffleContestCreate, Contest): c = self.end_condition if c.target_entry_amount and self.current_amount >= c.target_entry_amount: return True - if c.ends_at and datetime.now(tz=UTC) >= c.ends_at: - return True - return False + return bool(c.ends_at and datetime.now(tz=UTC) >= c.ends_at) def model_dump_mysql(self) -> dict[str, Any]: d = super().model_dump_mysql() @@ -213,7 +211,7 @@ class RaffleContest(RaffleContestCreate, Contest): return d @classmethod - def model_validate_mysql(cls, data: dict) -> Self: + def model_validate_mysql(cls, data: dict[str, Any]) -> Self: data["entry_rule"] = ContestEntryRule.model_validate(data["entry_rule"]) return super().model_validate_mysql(data) diff --git a/generalresearch/models/thl/demographics.py b/generalresearch/models/thl/demographics.py index b6a8be1..c11f8b2 100644 --- a/generalresearch/models/thl/demographics.py +++ b/generalresearch/models/thl/demographics.py @@ -76,7 +76,7 @@ class AgeGroup(Enum): return self.label -def calculate_demographic_metrics(opps: list[MarketplaceTask]) -> list: +def calculate_demographic_metrics(opps: list[MarketplaceTask]) -> list[dict[str, Any]]: """ Measurement: marketplace_survey_demographics tags: source (marketplace) @@ -86,7 +86,7 @@ def calculate_demographic_metrics(opps: list[MarketplaceTask]) -> list: """ source = {opp.source for opp in opps} assert len(source) == 1 - source = list(source)[0] + source = next(iter(source)) survey_cpi = defaultdict(list) target_open = defaultdict(int) for opp in opps: @@ -100,7 +100,7 @@ def calculate_demographic_metrics(opps: list[MarketplaceTask]) -> list: survey_counter = {k: len(v) for k, v in survey_cpi.items()} survey_counter = {k: {"count": v} for k, v in survey_counter.items() if v} - grp_stats = dict() + grp_stats = {} for grp, costs in survey_cpi.items(): stats = { "cost_min": np.min(costs), @@ -155,7 +155,7 @@ def calculate_used_question_metrics( """ source = {opp.source for opp in opps} assert len(source) == 1 - source = list(source)[0] + source = next(iter(source)) country_q_counter = defaultdict(Counter) for opp in opps: for q in opp.used_question_ids: diff --git a/generalresearch/models/thl/finance.py b/generalresearch/models/thl/finance.py index 8c94390..79a74a7 100644 --- a/generalresearch/models/thl/finance.py +++ b/generalresearch/models/thl/finance.py @@ -557,7 +557,7 @@ class BusinessBalances(BaseModel): they all explicitly are set """ - if any([pb.product_id is None for pb in v]): + if any(pb.product_id is None for pb in v): raise ValueError("'product_id' must be set for BusinessBalance children.") return v diff --git a/generalresearch/models/thl/ledger.py b/generalresearch/models/thl/ledger.py index 19dde20..a9fbbb1 100644 --- a/generalresearch/models/thl/ledger.py +++ b/generalresearch/models/thl/ledger.py @@ -1,7 +1,7 @@ from __future__ import annotations from datetime import UTC, datetime -from enum import StrEnum +from enum import IntEnum, StrEnum from typing import Annotated, Any, Literal, Self from uuid import uuid4 @@ -36,7 +36,7 @@ from generalresearch.models.thl.payout_format import ( from generalresearch.utils.enum import ReprEnumMeta -class Direction(int, Enum, metaclass=ReprEnumMeta): +class Direction(IntEnum, metaclass=ReprEnumMeta): """Entries on the debit side will increase debit normal accounts, while entries on the credit side will decrease them. Conversely, entries on the credit side will increase credit normal accounts, while entries on diff --git a/generalresearch/models/thl/offerwall/__init__.py b/generalresearch/models/thl/offerwall/__init__.py index e7c8e03..0c3d51d 100644 --- a/generalresearch/models/thl/offerwall/__init__.py +++ b/generalresearch/models/thl/offerwall/__init__.py @@ -267,7 +267,7 @@ class OfferWallRequest(BaseModel): # We need this so thl-core can refresh an offerwall in order to continue # a session d = self.model_dump(mode="json") - kwargs = dict() + kwargs = {} keys = [ "n_bins", "min_bin_size", diff --git a/generalresearch/models/thl/offerwall/base.py b/generalresearch/models/thl/offerwall/base.py index 33489df..33b9847 100644 --- a/generalresearch/models/thl/offerwall/base.py +++ b/generalresearch/models/thl/offerwall/base.py @@ -398,8 +398,10 @@ class OfferwallBucket(BaseModel): ) uri: HttpsUrl | None = Field( examples=[ - "https://task.generalresearch.com/api/v1/52d3f63b2709/797df4136c604a6c8599818296aae6d1/?i" - "=5ba2fe5010cc4d078fc3cc0b0cc264c3&b=test&66482fb=e7baf5e" + ( + "https://task.generalresearch.com/api/v1/52d3f63b2709/797df4136c604a6c8599818296aae6d1/?i" + "=5ba2fe5010cc4d078fc3cc0b0cc264c3&b=test&66482fb=e7baf5e" + ) ], description="The URL to send a respondent into. Must not edit this URL in any way", default=None, diff --git a/generalresearch/models/thl/payout_format.py b/generalresearch/models/thl/payout_format.py index 4d616b6..d29c9de 100644 --- a/generalresearch/models/thl/payout_format.py +++ b/generalresearch/models/thl/payout_format.py @@ -70,7 +70,7 @@ def format_payout_format(payout_format: str, payout_int: int) -> str: except TypeError: # "{payout()*1:}" - TypeError: 'int' object is not callable raise ValueError("Invalid type reference.") - except Exception: + except Exception: # noqa raise ValueError("Invalid payout transformation") formatstr = f"{{:{formatstr}}}" diff --git a/generalresearch/models/thl/product.py b/generalresearch/models/thl/product.py index a7ecd55..65ed177 100644 --- a/generalresearch/models/thl/product.py +++ b/generalresearch/models/thl/product.py @@ -957,7 +957,7 @@ class Product(BaseModel, validate_assignment=True): @field_validator("harmonizer_domain", mode="before") def harmonizer_domain_https(cls, s: str | None): # in the db, this has no scheme. accept both with a default of https:// - if s is not None and not (s.startswith("https://") or s.startswith("http://")): + if s is not None and not (s.startswith(("https://", "http://"))): s = f"https://{s}" return s @@ -1371,7 +1371,7 @@ class Product(BaseModel, validate_assignment=True): if self.payout_config.payout_transformation is None: return None payout_xform_func = self.get_payout_transformation_func() - kwargs = dict() + kwargs = {} if "user_wallet_balance" in inspect.signature(payout_xform_func).parameters: kwargs["user_wallet_balance"] = user_wallet_balance user_payout: Decimal = payout_xform_func(bp_payout, **kwargs) diff --git a/generalresearch/models/thl/profiling/marketplace.py b/generalresearch/models/thl/profiling/marketplace.py index ad4ce80..0129e38 100644 --- a/generalresearch/models/thl/profiling/marketplace.py +++ b/generalresearch/models/thl/profiling/marketplace.py @@ -82,10 +82,9 @@ class MarketplaceQuestion(BaseModel, ABC): # question has more than 6. repr_args = list(self.__repr_args__()) for n, (k, v) in enumerate(repr_args): - if k == "options": - if v and len(v) > 6: - v = v[:3] + ["..."] + v[-3:] - repr_args[n] = ("options", v) + if k == "options" and v and len(v) > 6: + v = v[:3] + ["..."] + v[-3:] + repr_args[n] = ("options", 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 diff --git a/generalresearch/models/thl/report_task.py b/generalresearch/models/thl/report_task.py index d29599d..299ba90 100644 --- a/generalresearch/models/thl/report_task.py +++ b/generalresearch/models/thl/report_task.py @@ -28,7 +28,7 @@ def prioritize_report_values( return None report_values = list(set(report_values)) random.shuffle(report_values) - return sorted(report_values, key=lambda x: REPORT_PRIORITY[x])[-1] + return max(report_values, key=lambda x: REPORT_PRIORITY[x]) class ReportTask(BaseModel): diff --git a/generalresearch/models/thl/session.py b/generalresearch/models/thl/session.py index fe7194a..e4e264f 100644 --- a/generalresearch/models/thl/session.py +++ b/generalresearch/models/thl/session.py @@ -234,10 +234,7 @@ class WallBase(BaseModel): return self.is_visible() and self.status == Status.COMPLETE def allow_session(self) -> bool: - if self.status == Status.COMPLETE: - return False - - return True + return self.status != Status.COMPLETE def update(self, **kwargs) -> None: """ @@ -969,10 +966,7 @@ class Session(BaseModel): return True # Hard limit of 40 wall events per session - if len(self.wall_events) >= 40: - return True - - return False + return len(self.wall_events) >= 40 def determine_payments( self, @@ -985,6 +979,7 @@ class Session(BaseModel): ) product = self.user.product + assert product # Handle brokerage product payouts bp_pay: Decimal = product.determine_bp_payment(thl_net) commission_amount: Decimal = thl_net - bp_pay diff --git a/generalresearch/models/thl/soft_pair.py b/generalresearch/models/thl/soft_pair.py index f3b2b6f..7c2f36e 100644 --- a/generalresearch/models/thl/soft_pair.py +++ b/generalresearch/models/thl/soft_pair.py @@ -50,7 +50,7 @@ class SoftPairResult: return ( self.survey_id + ":" - + ";".join(sorted(set([c.question_id for c in self.conditions]))) + + ";".join(sorted({c.question_id for c in self.conditions})) ) else: return None diff --git a/generalresearch/models/thl/survey/penalty.py b/generalresearch/models/thl/survey/penalty.py index 04f8e20..755d25c 100644 --- a/generalresearch/models/thl/survey/penalty.py +++ b/generalresearch/models/thl/survey/penalty.py @@ -56,7 +56,7 @@ class TeamSurveyPenalty(SurveyPenalty): Penalty = Annotated[ - Union[BPSurveyPenalty, TeamSurveyPenalty], + BPSurveyPenalty | TeamSurveyPenalty, Field(discriminator="kind"), ] PenaltyListAdapter = TypeAdapter(list[Penalty]) diff --git a/generalresearch/models/thl/survey/task_collection.py b/generalresearch/models/thl/survey/task_collection.py index b80166f..d8db0d1 100644 --- a/generalresearch/models/thl/survey/task_collection.py +++ b/generalresearch/models/thl/survey/task_collection.py @@ -38,7 +38,8 @@ class TaskCollection(BaseModel): except pa.errors.SchemaErrors as exc: idx = exc.failure_cases["index"] if len(idx) >= len(df) * 0.10: - raise exc + raise + logger.info(f"{self.__repr_name__()}:handle_df:{json.dumps(exc.message)}") df.drop(index=list(idx), inplace=True) # we need to redo the validation after removing failing rows! diff --git a/generalresearch/models/thl/task_status.py b/generalresearch/models/thl/task_status.py index de767d6..7719b18 100644 --- a/generalresearch/models/thl/task_status.py +++ b/generalresearch/models/thl/task_status.py @@ -224,11 +224,12 @@ class TaskStatusResponse(BaseModel): return v or 0 @field_validator("kwargs", mode="after") - def sanitize_kwargs(cls, v: dict | None) -> dict | None: + def sanitize_kwargs(cls, v: dict[str, Any] | None) -> dict[str, Any] | None: if v and "clicked_timestamp" in v: try: - clicked_timestamp = datetime.strptime( - v["clicked_timestamp"], "%Y-%m-%d %H:%M:%S.%f" + clicked_timestamp = datetime.strptime( # noqa + date_string=v["clicked_timestamp"], + format="%Y-%m-%d %H:%M:%S.%f", ) v["clicked_timestamp"] = ( clicked_timestamp.isoformat(timespec="microseconds") + "Z" @@ -238,7 +239,7 @@ class TaskStatusResponse(BaseModel): return v @model_validator(mode="before") - def transform_user_payout(cls, d): + def transform_user_payout(cls, d: dict[str, Any]): # If the user_payout is None and there is a payout_format, make the user_payout 0 if d.get("user_payout") is None and d.get("payout_format"): d["user_payout"] = 0 diff --git a/generalresearch/pg_helper.py b/generalresearch/pg_helper.py index 1d5d30b..a397247 100644 --- a/generalresearch/pg_helper.py +++ b/generalresearch/pg_helper.py @@ -108,10 +108,8 @@ class PostgresConfig: def execute_write(self, query, params=None) -> int: cmd = query.lstrip().upper() - assert ( - cmd.startswith("INSERT") - or cmd.startswith("UPDATE") - or cmd.startswith("DELETE") + assert cmd.startswith( + ("INSERT", "UPDATE", "DELETE") ), "Supports INSERT/UPDATE only" with self.make_connection() as conn: diff --git a/generalresearch/sql_helper.py b/generalresearch/sql_helper.py index 08b660d..ae2b8d8 100644 --- a/generalresearch/sql_helper.py +++ b/generalresearch/sql_helper.py @@ -315,7 +315,7 @@ class SqlHelper(SqlConnector): field_names = ["`" + x + "`" for x in field_names] field_name_str = ",".join(field_names) if filter_d: - lookup_vals = " AND ".join([f"`{fn}`=%({fn})s" for fn in filter_d.keys()]) + lookup_vals = " AND ".join([f"`{fn}`=%({fn})s" for fn in filter_d]) lookup_str = f" WHERE {lookup_vals}" else: lookup_str = "" diff --git a/generalresearch/utils/grpc_logger.py b/generalresearch/utils/grpc_logger.py index 59f7471..8f2f454 100644 --- a/generalresearch/utils/grpc_logger.py +++ b/generalresearch/utils/grpc_logger.py @@ -33,9 +33,11 @@ try: response = handler_func(request, context) code = context.code() or grpc.StatusCode.OK return response - except Exception as e: + + except Exception: code = context.code() or grpc.StatusCode.INTERNAL - raise e + raise + finally: duration_ms = int((time.time() - start_time) * 1000) peer = context.peer() or "unknown" diff --git a/generalresearch/wall_status_codes/fullcircle.py b/generalresearch/wall_status_codes/fullcircle.py index aeaa4c7..cd9fdff 100644 --- a/generalresearch/wall_status_codes/fullcircle.py +++ b/generalresearch/wall_status_codes/fullcircle.py @@ -29,7 +29,7 @@ status_codes_ext_map: dict[StatusCode1, list[str]] = { StatusCode1.PS_FAIL: [], StatusCode1.PS_OVERQUOTA: [], } -ext_status_code_map: dict[str, StatusCode1] = dict() +ext_status_code_map: dict[str, StatusCode1] = {} for k, v in status_codes_ext_map.items(): k: StatusCode1 v: list[str] diff --git a/generalresearch/wall_status_codes/innovate.py b/generalresearch/wall_status_codes/innovate.py index 936ee6c..e3d2468 100644 --- a/generalresearch/wall_status_codes/innovate.py +++ b/generalresearch/wall_status_codes/innovate.py @@ -38,7 +38,7 @@ status_codes_ext_map: dict[StatusCode1, list[str]] = { StatusCode1.PS_FAIL: ["5"], StatusCode1.PS_OVERQUOTA: ["7"], } -ext_status_code_map = dict() +ext_status_code_map = {} for k, v in status_codes_ext_map.items(): for vv in v: ext_status_code_map[status_codes_ext_map.get(vv, vv)] = k diff --git a/generalresearch/wall_status_codes/lucid.py b/generalresearch/wall_status_codes/lucid.py index 3cc1b5e..c4c5e90 100644 --- a/generalresearch/wall_status_codes/lucid.py +++ b/generalresearch/wall_status_codes/lucid.py @@ -102,7 +102,7 @@ status_codes_ext_map: dict[StatusCode1, list[str]] = { StatusCode1.PS_OVERQUOTA: ["40", "41", "42"], } -ext_status_code_map: dict[str, StatusCode1] = dict() +ext_status_code_map: dict[str, StatusCode1] = {} for k, v in status_codes_ext_map.items(): k: StatusCode1 v: list[str] diff --git a/generalresearch/wall_status_codes/morning.py b/generalresearch/wall_status_codes/morning.py index ffa6be2..6f63b82 100644 --- a/generalresearch/wall_status_codes/morning.py +++ b/generalresearch/wall_status_codes/morning.py @@ -97,7 +97,7 @@ status_codes_ext_map: dict[StatusCode1, list[str]] = { "quota_invalid_for_bid", ], } -ext_status_code_map: dict[str, StatusCode1] = dict() +ext_status_code_map: dict[str, StatusCode1] = {} for k, v in status_codes_ext_map.items(): k: StatusCode1 v: list[str] diff --git a/generalresearch/wall_status_codes/pollfish.py b/generalresearch/wall_status_codes/pollfish.py index a5c6e25..e1ad12a 100644 --- a/generalresearch/wall_status_codes/pollfish.py +++ b/generalresearch/wall_status_codes/pollfish.py @@ -58,7 +58,7 @@ status_codes_ext_map: dict[StatusCode1, list[str]] = { ], StatusCode1.PS_OVERQUOTA: ["quota_full", "survey_closed", "survey_expired"], } -ext_status_code_map = dict() +ext_status_code_map = {} for k, v in status_codes_ext_map.items(): k: StatusCode1 v: list[str] diff --git a/test_utils/managers/contest/conftest.py b/test_utils/managers/contest/conftest.py index 67935e7..a9375f6 100644 --- a/test_utils/managers/contest/conftest.py +++ b/test_utils/managers/contest/conftest.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import pytest from generalresearch.managers.base import Permission @@ -11,8 +13,6 @@ def contest_manager(thl_web_rw: PostgresConfig) -> ContestManager: assert thl_web_rw.dsn.path assert "/unittest-" in thl_web_rw.dsn.path - from generalresearch.managers.thl.contest_manager import ContestManager - return ContestManager( pg_config=thl_web_rw, permissions=[ diff --git a/test_utils/models/contest/conftest.py b/test_utils/models/contest/conftest.py index d9c8a6b..84930b8 100644 --- a/test_utils/models/contest/conftest.py +++ b/test_utils/models/contest/conftest.py @@ -38,6 +38,26 @@ from generalresearch.models.thl.user import User # === Managers === +# --- Factories --- + + +@pytest.fixture(scope="function") +def raffle_contest_factory( + product_user_wallet_yes: Product, + raffle_contest_create: RaffleContestCreate, + contest_manager: ContestManager, +) -> Callable[..., RaffleContest]: + + def _inner(**kwargs): + raffle_contest_create.update(**kwargs) + return contest_manager.create( + product_id=product_user_wallet_yes.uuid, + contest_create=raffle_contest_create, + ) + + return _inner + + # === Models === @@ -82,23 +102,6 @@ def raffle_contest( ) -@pytest.fixture(scope="function") -def raffle_contest_factory( - product_user_wallet_yes: Product, - raffle_contest_create: RaffleContestCreate, - contest_manager: ContestManager, -) -> Callable[..., RaffleContest]: - - def _inner(**kwargs): - raffle_contest_create.update(**kwargs) - return contest_manager.create( - product_id=product_user_wallet_yes.uuid, - contest_create=raffle_contest_create, - ) - - return _inner - - @pytest.fixture def milestone_contest_create() -> MilestoneContestCreate: from generalresearch.models.thl.contest import ( diff --git a/test_utils/spectrum/conftest.py b/test_utils/spectrum/conftest.py index 7cd9321..a8ce9d9 100644 --- a/test_utils/spectrum/conftest.py +++ b/test_utils/spectrum/conftest.py @@ -1,32 +1,32 @@ from __future__ import annotations -import logging import time from datetime import UTC, datetime from decimal import Decimal -from typing import TYPE_CHECKING, Any +from typing import Any import pytest +from generalresearch.config import GRLBaseSettings from generalresearch.managers.spectrum.survey import ( SpectrumCriteriaManager, SpectrumSurveyManager, ) -from generalresearch.models.spectrum.survey import SpectrumSurvey +from generalresearch.models import ( + LogicalOperator, +) +from generalresearch.models.spectrum.survey import ( + SpectrumCondition, + SpectrumSurvey, +) +from generalresearch.models.thl.survey.condition import ConditionValueType from generalresearch.sql_helper import SqlHelper -from .surveys_json import CONDITIONS, SURVEYS_JSON - -if TYPE_CHECKING: - from generalresearch.config import GRLBaseSettings - @pytest.fixture(scope="session") def spectrum_rw(settings: GRLBaseSettings) -> SqlHelper: - logging.info(f"{settings.spectrum_rw_db=}") - assert settings.spectrum_rw_db is not None - assert "/unittest-" in settings.spectrum_rw_db.path + assert "/unittest-" in str(settings.spectrum_rw_db.path) return SqlHelper( dsn=settings.spectrum_rw_db, @@ -38,27 +38,36 @@ def spectrum_rw(settings: GRLBaseSettings) -> SqlHelper: @pytest.fixture(scope="session") def spectrum_criteria_manager(spectrum_rw: SqlHelper) -> SpectrumCriteriaManager: + assert spectrum_rw.dsn + assert spectrum_rw.dsn.path assert "/unittest-" in spectrum_rw.dsn.path return SpectrumCriteriaManager(spectrum_rw) @pytest.fixture(scope="session") def spectrum_survey_manager(spectrum_rw: SqlHelper) -> SpectrumSurveyManager: + assert spectrum_rw.dsn + assert spectrum_rw.dsn.path assert "/unittest-" in spectrum_rw.dsn.path return SpectrumSurveyManager(spectrum_rw) @pytest.fixture(scope="session") def setup_spectrum_surveys( - spectrum_rw: SqlHelper, spectrum_survey_manager, spectrum_criteria_manager + spectrum_rw: SqlHelper, + spectrum_survey_manager: SpectrumSurveyManager, + spectrum_criteria_manager: SpectrumCriteriaManager, + spectrum_conditions: list[SpectrumCondition], + spectrum_api_surveys_json: list[str], ) -> None: now = datetime.now(UTC) # make sure these example surveys exist in db - surveys = [SpectrumSurvey.model_validate_json(x) for x in SURVEYS_JSON] + surveys = [SpectrumSurvey.model_validate_json(x) for x in spectrum_api_surveys_json] for s in surveys: s.modified_api = datetime.now(tz=UTC) + spectrum_survey_manager.create_or_update(surveys) - spectrum_criteria_manager.update(CONDITIONS) + spectrum_criteria_manager.update(spectrum_conditions) # and make sure they have allocation for 687 spectrum_rw.execute_sql_query( @@ -198,6 +207,46 @@ def spectrum_api_surveys_json() -> list[str]: ] +def spectrum_conditions() -> list[SpectrumCondition]: + # 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, + ) + return [c1, c2, c3, c4, c5] + + @pytest.fixture(scope="session") def spectrum_api_survey_json() -> dict[str, Any]: return { diff --git a/tests/incite/collections/test_df_collection_base.py b/tests/incite/collections/test_df_collection_base.py index b9f0181..c236700 100644 --- a/tests/incite/collections/test_df_collection_base.py +++ b/tests/incite/collections/test_df_collection_base.py @@ -9,6 +9,7 @@ from generalresearch.incite.collections import ( DFCollection, DFCollectionType, ) +from generalresearch.pg_helper import PostgresConfig if TYPE_CHECKING: from generalresearch.incite.base import GRLDatasets @@ -45,7 +46,9 @@ class TestDFCollectionBase: class TestDFCollectionBaseProperties: @pytest.mark.skip - def test_df_collection_items(self, mnt_filepath: GRLDatasets, df_coll_type): + def test_df_collection_items( + self, mnt_filepath: GRLDatasets, df_coll_type: DFCollectionType + ): instance = DFCollection( data_type=df_coll_type, start=datetime(year=1800, month=1, day=1, tzinfo=UTC), @@ -57,7 +60,9 @@ class TestDFCollectionBaseProperties: assert len(instance.interval_range) == len(instance.items) assert len(instance.items) == 366 - def test_df_collection_progress(self, mnt_filepath: GRLDatasets, df_coll_type): + def test_df_collection_progress( + self, mnt_filepath: GRLDatasets, df_coll_type: DFCollectionType + ): instance = DFCollection( data_type=df_coll_type, start=datetime(year=1800, month=1, day=1, tzinfo=UTC), @@ -70,7 +75,9 @@ class TestDFCollectionBaseProperties: assert isinstance(instance.progress, pd.DataFrame) assert instance.progress.shape == (366, 6) - def test_df_collection_schema(self, mnt_filepath: GRLDatasets, df_coll_type): + def test_df_collection_schema( + self, mnt_filepath: GRLDatasets, df_coll_type: DFCollectionType + ): instance1 = DFCollection( data_type=DFCollectionType.WALL, archive_path=mnt_filepath.data_src ) @@ -87,9 +94,9 @@ class TestDFCollectionBaseProperties: class TestDFCollectionBaseMethods: @pytest.mark.skip - def test_initial_load(self, mnt_filepath: GRLDatasets, thl_web_rr): + def test_initial_load(self, mnt_filepath: GRLDatasets, thl_web_rr: PostgresConfig): instance = DFCollection( - pg_config=thl_web_rr: PostgresConfig, + pg_config=thl_web_rr, data_type=DFCollectionType.USER, start=datetime(year=2022, month=1, day=1, minute=0, tzinfo=UTC), finished=datetime(year=2022, month=1, day=1, minute=5, tzinfo=UTC), diff --git a/tests/incite/collections/test_df_collection_item_base.py b/tests/incite/collections/test_df_collection_item_base.py index 9a2ecf3..e0171c2 100644 --- a/tests/incite/collections/test_df_collection_item_base.py +++ b/tests/incite/collections/test_df_collection_item_base.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from datetime import UTC, datetime from typing import TYPE_CHECKING @@ -19,7 +21,7 @@ df_collection_types = [e for e in DFCollectionType if e is not DFCollectionType. @pytest.mark.parametrize("df_coll_type", df_collection_types) class TestDFCollectionItemBase: - def test_init(self, mnt_filepath: GRLDatasets, df_coll_type): + def test_init(self, mnt_filepath: GRLDatasets, df_coll_type: DFCollectionType): collection = DFCollection( data_type=df_coll_type, offset="100d", @@ -38,14 +40,16 @@ class TestDFCollectionItemBase: class TestDFCollectionItemProperties: @pytest.mark.skip - def test_filename(self, df_coll_type): + def test_filename(self, df_coll_type: DFCollectionType): pass @pytest.mark.parametrize("df_coll_type", df_collection_types) class TestDFCollectionItemMethods: - def test_has_mysql_false(self, mnt_filepath: GRLDatasets, df_coll_type): + def test_has_mysql_false( + self, mnt_filepath: GRLDatasets, df_coll_type: DFCollectionType + ): collection = DFCollection( data_type=df_coll_type, offset="100d", @@ -58,7 +62,10 @@ class TestDFCollectionItemMethods: assert not instance1.has_mysql() def test_has_mysql_true( - self, thl_web_rr: PostgresConfig, mnt_filepath: GRLDatasets, df_coll_type + self, + thl_web_rr: PostgresConfig, + mnt_filepath: GRLDatasets, + df_coll_type: DFCollectionType, ): collection = DFCollection( data_type=df_coll_type, @@ -66,7 +73,7 @@ class TestDFCollectionItemMethods: start=datetime(year=1800, month=1, day=1, tzinfo=UTC), finished=datetime(year=1900, month=1, day=1, tzinfo=UTC), archive_path=mnt_filepath.archive_path(enum_type=df_coll_type), - pg_config=thl_web_rr: PostgresConfig, + pg_config=thl_web_rr, ) # Has RR, assume unittest server is online @@ -74,5 +81,5 @@ class TestDFCollectionItemMethods: assert instance2.has_mysql() @pytest.mark.skip - def test_update_partial_archive(self, df_coll_type): + def test_update_partial_archive(self, df_coll_type: DFCollectionType): pass diff --git a/tests/incite/collections/test_df_collection_thl_marketplaces.py b/tests/incite/collections/test_df_collection_thl_marketplaces.py index d2d3ce4..b4b5b00 100644 --- a/tests/incite/collections/test_df_collection_thl_marketplaces.py +++ b/tests/incite/collections/test_df_collection_thl_marketplaces.py @@ -4,6 +4,7 @@ from itertools import product import pytest from pandera.pandas import Column, DataFrameSchema, Index +from generalresearch.incite.base import GRLDatasets from generalresearch.incite.collections import DFCollection, DFCollectionType from generalresearch.incite.collections.thl_marketplaces import ( InnovateSurveyHistoryCollection, @@ -11,6 +12,7 @@ from generalresearch.incite.collections.thl_marketplaces import ( SagoSurveyHistoryCollection, SpectrumSurveyTimeseriesCollection, ) +from generalresearch.pg_helper import PostgresConfig def combo_object(): @@ -29,7 +31,13 @@ def combo_object(): @pytest.mark.parametrize("df_coll, offset", combo_object()) class TestDFCollection_thl_marketplaces: - def test_init(self, mnt_filepath, df_coll, offset, spectrum_rw): + def test_init( + self, + mnt_filepath: GRLDatasets, + df_coll: DFCollection, + offset: str, + spectrum_rw: PostgresConfig, + ): assert issubclass(df_coll, DFCollection) # This is stupid, but we need to pull the default from the @@ -38,7 +46,7 @@ class TestDFCollection_thl_marketplaces: assert isinstance(data_type, DFCollectionType) # (1) Can't be totally empty, needs a path... - with pytest.raises(expected_exception=Exception) as cm: + with pytest.raises(expected_exception=Exception): instance = df_coll() # (2) Confirm it only needs the archive_path @@ -61,7 +69,7 @@ class TestDFCollection_thl_marketplaces: assert isinstance(instance._schema, DataFrameSchema) assert isinstance(instance._schema.index, Index) - for c in instance._schema.columns.keys(): + for c in instance._schema.columns: assert isinstance(c, str) col = instance._schema.columns[c] assert isinstance(col, Column) diff --git a/tests/incite/collections/test_df_collection_thl_web.py b/tests/incite/collections/test_df_collection_thl_web.py index bcdeb83..6d509bc 100644 --- a/tests/incite/collections/test_df_collection_thl_web.py +++ b/tests/incite/collections/test_df_collection_thl_web.py @@ -3,19 +3,16 @@ from __future__ import annotations from collections.abc import Generator from datetime import datetime from itertools import product -from typing import TYPE_CHECKING import dask.dataframe as dd import pandas as pd import pytest from pandera.pandas import DataFrameSchema -from generalresearch.incite.collections import DFCollection, DFCollectionType - -if TYPE_CHECKING: - from generalresearch.incite.collections import ( - DFCollectionType, - ) +from generalresearch.incite.collections import ( + DFCollection, + DFCollectionType, +) def combo_object() -> Generator[tuple]: @@ -39,7 +36,10 @@ def combo_object() -> Generator[tuple]: class TestDFCollection_thl_web: def test_init( - self, df_collection_data_type: DFCollectionType, offset: str, df_collection + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, ): assert isinstance(df_collection_data_type, DFCollectionType) assert isinstance(df_collection, DFCollection) @@ -50,12 +50,12 @@ class TestDFCollection_thl_web: ) class TestDFCollection_thl_web_Properties: - def test_items(self, df_collection): + def test_items(self, df_collection: DFCollection): assert isinstance(df_collection.items, list) for i in df_collection.items: assert i._collection == df_collection - def test__schema(self, df_collection): + def test__schema(self, df_collection: DFCollection): assert isinstance(df_collection._schema, DataFrameSchema) @@ -65,16 +65,16 @@ class TestDFCollection_thl_web_Properties: class TestDFCollection_thl_web_BaseProperties: @pytest.mark.skip - def test__interval_range(self, df_collection): + def test__interval_range(self, df_collection: DFCollection): pass - def test_interval_start(self, df_collection): + def test_interval_start(self, df_collection: DFCollection): assert isinstance(df_collection.interval_start, datetime) - def test_interval_range(self, df_collection): + def test_interval_range(self, df_collection: DFCollection): assert isinstance(df_collection.interval_range, list) - def test_progress(self, df_collection): + def test_progress(self, df_collection: DFCollection): assert isinstance(df_collection.progress, pd.DataFrame) @@ -84,17 +84,21 @@ class TestDFCollection_thl_web_BaseProperties: class TestDFCollection_thl_web_Methods: @pytest.mark.skip - def test_initial_loads(self, df_collection_data_type, df_collection, offset): + def test_initial_loads( + self, df_collection_data_type, df_collection: DFCollection, offset: str + ): pass @pytest.mark.skip def test_fetch_force_rr_latest( - self, df_collection_data_type, df_collection, offset: str + self, df_collection_data_type, df_collection: DFCollection, offset: str ): pass @pytest.mark.skip - def test_force_rr_latest(self, df_collection_data_type, df_collection, offset): + def test_force_rr_latest( + self, df_collection_data_type, df_collection: DFCollection, offset: str + ): pass @@ -103,63 +107,108 @@ class TestDFCollection_thl_web_Methods: ) class TestDFCollection_thl_web_BaseMethods: - def test_fetch_all_paths(self, df_collection_data_type, offset: str, df_collection): + def test_fetch_all_paths( + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, + ): res = df_collection.fetch_all_paths( items=None, force_rr_latest=False, include_partial=False ) assert isinstance(res, list) @pytest.mark.skip - def test_ddf(self, df_collection_data_type, offset: str, df_collection): + def test_ddf( + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, + ): res = df_collection.ddf() assert isinstance(res, dd.DataFrame) # -- cleanup -- @pytest.mark.skip def test_schedule_cleanup( - self, df_collection_data_type, offset: str, df_collection + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, ): pass @pytest.mark.skip - def test_cleanup(self, df_collection_data_type, offset: str, df_collection): + def test_cleanup( + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, + ): pass @pytest.mark.skip def test_cleanup_partials( - self, df_collection_data_type, offset: str, df_collection + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, ): pass @pytest.mark.skip def test_clear_tmp_archives( - self, df_collection_data_type, offset: str, df_collection + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, ): pass @pytest.mark.skip def test_clear_corrupt_archives( - self, df_collection_data_type, offset: str, df_collection + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, ): pass @pytest.mark.skip def test_rebuild_symlinks( - self, df_collection_data_type, offset: str, df_collection + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, ): pass # -- Source timing -- @pytest.mark.skip - def test_get_item(self, df_collection_data_type, offset: str, df_collection): + def test_get_item( + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, + ): pass @pytest.mark.skip - def test_get_item_start(self, df_collection_data_type, offset: str, df_collection): + def test_get_item_start( + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, + ): pass @pytest.mark.skip - def test_get_items(self, df_collection_data_type, offset: str, df_collection): + def test_get_items( + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, + ): # If we get all the items from the start of the collection, it # should include all the items! res1 = df_collection.items @@ -168,18 +217,27 @@ class TestDFCollection_thl_web_BaseMethods: @pytest.mark.skip def test_get_items_from_year( - self, df_collection_data_type, offset: str, df_collection + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, ): pass @pytest.mark.skip def test_get_items_last90( - self, df_collection_data_type, offset: str, df_collection + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, ): pass @pytest.mark.skip def test_get_items_last365( - self, df_collection_data_type, offset: str, df_collection + self, + df_collection_data_type: DFCollectionType, + offset: str, + df_collection: DFCollection, ): pass diff --git a/tests/incite/mergers/foundations/test_user_id_product.py b/tests/incite/mergers/foundations/test_user_id_product.py index 10802e5..7367056 100644 --- a/tests/incite/mergers/foundations/test_user_id_product.py +++ b/tests/incite/mergers/foundations/test_user_id_product.py @@ -1,11 +1,15 @@ +from __future__ import annotations + from datetime import UTC, datetime, timedelta from itertools import product import pandas as pd import pytest +from dask.distributed import Client as DaskClient # noinspection PyUnresolvedReferences from generalresearch.incite.mergers.foundations.user_id_product import ( + UserIdProductMerge, UserIdProductMergeItem, ) @@ -23,14 +27,21 @@ from generalresearch.incite.mergers.foundations.user_id_product import ( class TestUserIDProduct: @pytest.mark.skip - def test_base(self, client_no_amm, user_id_product_merge): + def test_base( + self, client_no_amm: DaskClient, user_id_product_merge: UserIdProductMerge + ): ddf = user_id_product_merge.ddf() df = client_no_amm.compute(collections=ddf, sync=True) assert isinstance(df, pd.DataFrame) assert not df.empty @pytest.mark.skip - def test_base_item(self, client_no_amm, user_id_product_merge, user_collection): + def test_base_item( + self, + client_no_amm: DaskClient, + user_id_product_merge: UserIdProductMerge, + user_collection, + ): assert len(user_id_product_merge.items) == 1 for item in user_id_product_merge.items: @@ -40,7 +51,7 @@ class TestUserIDProduct: try: modified_time1 = path.stat().st_mtime - except Exception: + except OSError: modified_time1 = 0 user_id_product_merge.build(client=client_no_amm, user_coll=user_collection) @@ -49,7 +60,9 @@ class TestUserIDProduct: assert modified_time2 > modified_time1 @pytest.mark.skip - def test_read(self, client_no_amm, user_id_product_merge): + def test_read( + self, client_no_amm: DaskClient, user_id_product_merge: UserIdProductMerge + ): users_ddf = user_id_product_merge.ddf() df = client_no_amm.compute(collections=users_ddf, sync=True) diff --git a/tests/incite/mergers/test_pop_ledger.py b/tests/incite/mergers/test_pop_ledger.py index 529a641..2146344 100644 --- a/tests/incite/mergers/test_pop_ledger.py +++ b/tests/incite/mergers/test_pop_ledger.py @@ -86,9 +86,7 @@ class TestMergePOPLedger: # -- - user_wallet_account: LedgerAccount = ( - thl_ledger_manager.get_account_or_create_user_wallet(user=u) - ) + 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() @@ -295,7 +293,7 @@ class TestMergePOPLedger: assert isinstance(df.index, pd.Index) assert isinstance(df.index, pd.DatetimeIndex) - bp_account_balance = thl_ledger_manager.get_account_balance(account=bp_account) + thl_ledger_manager.get_account_balance(account=bp_account) # Initial sum initial_sum = df.sum().sum() diff --git a/tests/incite/test_collection_base.py b/tests/incite/test_collection_base.py index d6ce2b1..577eda9 100644 --- a/tests/incite/test_collection_base.py +++ b/tests/incite/test_collection_base.py @@ -241,7 +241,7 @@ class TestCollectionBaseMethodsCleanup: assert "Must override" in str(cm.value) -class TestCollectionBaseMethodsCleanup: +class TestCollectionBaseMethodsCleanup2: @pytest.mark.skip def test_cleanup_partials(self, mnt_filepath: GRLDatasets): diff --git a/tests/managers/thl/test_contest/test_milestone.py b/tests/managers/thl/test_contest/test_milestone.py index e3889bc..aa738e1 100644 --- a/tests/managers/thl/test_contest/test_milestone.py +++ b/tests/managers/thl/test_contest/test_milestone.py @@ -15,6 +15,7 @@ from generalresearch.models.thl.contest.milestone import ( MilestoneContestCreate, MilestoneUserView, ) +from generalresearch.models.thl.contest.raffle import RaffleContest from generalresearch.models.thl.product import Product from generalresearch.models.thl.user import User @@ -241,7 +242,7 @@ class TestMilestoneContestUserViews: def test_list_user_eligible_country( self, user_with_wallet: User, - contest_factory: Callable[..., Contest], + raffle_contest_factory: Callable[..., Contest], thl_ledger_manager: ThlLedgerManager, contest_manager: ContestManager, ): @@ -252,7 +253,7 @@ class TestMilestoneContestUserViews: assert len(cs) == 0 # Create a contest. It'll be in the US/CA - contest_factory(country_isos={"us", "ca"}) + raffle_contest_factory(country_isos={"us", "ca"}) # Not eligible in mexico cs = contest_manager.get_many_by_user_eligible( @@ -265,7 +266,7 @@ class TestMilestoneContestUserViews: assert len(cs) == 1 # Create another, any country - contest_factory(country_isos=None) + raffle_contest_factory(country_isos=None) cs = contest_manager.get_many_by_user_eligible( user=user_with_wallet, country_iso="mx" ) @@ -278,12 +279,12 @@ class TestMilestoneContestUserViews: def test_list_user_eligible( self, user_with_money: User, - contest_factory: Callable[..., Contest], + raffle_contest_factory: Callable[..., RaffleContest], thl_ledger_manager: ThlLedgerManager, contest_manager: ContestManager, ): # User reaches milestone after 1 complete - c = contest_factory(target_amount=1) + c = raffle_contest_factory(target_amount=1) user = user_with_money cs = contest_manager.get_many_by_user_eligible( diff --git a/tests/managers/thl/test_ledger/test_thl_lm_tx__user_payouts.py b/tests/managers/thl/test_ledger/test_thl_lm_tx__user_payouts.py index 5fb6935..82dc143 100644 --- a/tests/managers/thl/test_ledger/test_thl_lm_tx__user_payouts.py +++ b/tests/managers/thl/test_ledger/test_thl_lm_tx__user_payouts.py @@ -277,11 +277,12 @@ class TestLedgerManagerAMT: thl_ledger_manager.create_tx_user_payout_cancelled( user, payout_event=pe, skip_flag_check=True ) - with pytest.raises(expected_exception=LedgerTransactionConditionFailedError): - with caplog.at_level(logging.WARNING): - thl_ledger_manager.create_tx_user_payout_complete( - user, payout_event=pe, skip_flag_check=True - ) + with pytest.raises( + expected_exception=LedgerTransactionConditionFailedError + ), caplog.at_level(logging.WARNING): + thl_ledger_manager.create_tx_user_payout_complete( + user, payout_event=pe, skip_flag_check=True + ) assert "trying to complete payout that was already cancelled" in caplog.text cash = thl_ledger_manager.get_account_cash() @@ -319,11 +320,12 @@ class TestLedgerManagerAMT: thl_ledger_manager.create_tx_user_payout_complete( user, payout_event=pe2, skip_flag_check=True ) - with pytest.raises(expected_exception=LedgerTransactionConditionFailedError): - with caplog.at_level(logging.WARNING): - thl_ledger_manager.create_tx_user_payout_cancelled( - user, payout_event=pe2, skip_flag_check=True - ) + with pytest.raises( + expected_exception=LedgerTransactionConditionFailedError + ), caplog.at_level(logging.WARNING): + thl_ledger_manager.create_tx_user_payout_cancelled( + user, payout_event=pe2, skip_flag_check=True + ) assert "trying to cancel payout that was already completed" in caplog.text diff --git a/tests/models/spectrum/test_survey.py b/tests/models/spectrum/test_survey.py index f97860f..bad6857 100644 --- a/tests/models/spectrum/test_survey.py +++ b/tests/models/spectrum/test_survey.py @@ -407,44 +407,12 @@ class TestSpectrumSurvey: ) -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] +def test_spectrum_something( + spectrum_conditions: list[SpectrumCondition], spectrum_api_surveys_json: list[str] +): + + c1 = spectrum_conditions[0] + c3 = spectrum_conditions[2] survey = SpectrumSurvey.model_validate_json(spectrum_api_surveys_json[0]) assert c1.criterion_hash in survey.qualifications -- cgit v1.2.3